diff --git a/mindspore/core/abstract/prim_arrays.cc b/mindspore/core/abstract/prim_arrays.cc index af1bb853da..ba65e6adc2 100644 --- a/mindspore/core/abstract/prim_arrays.cc +++ b/mindspore/core/abstract/prim_arrays.cc @@ -24,6 +24,7 @@ #include "utils/shape_utils.h" #include "ops/op_utils.h" #include "utils/anf_utils.h" +#include "utils/check_convert_utils.h" namespace mindspore { namespace abstract { diff --git a/mindspore/core/abstract/primitive_infer_map.cc b/mindspore/core/abstract/primitive_infer_map.cc index 09adc4911e..11c414af56 100644 --- a/mindspore/core/abstract/primitive_infer_map.cc +++ b/mindspore/core/abstract/primitive_infer_map.cc @@ -16,6 +16,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "abstract/primitive_infer_map.h" #include #include diff --git a/mindspore/core/abstract/primitive_infer_map.h b/mindspore/core/abstract/primitive_infer_map.h index 91a9494f74..0ac61b81e6 100644 --- a/mindspore/core/abstract/primitive_infer_map.h +++ b/mindspore/core/abstract/primitive_infer_map.h @@ -69,8 +69,8 @@ class RegisterStandardPrimitiveEvalHelper { static auto helper_##name = \ abstract::RegisterStandardPrimitiveEvalHelper(primitive, infer_impl, infer_value_impl, is_white_list); \ std::shared_ptr GetDefaultPrimC##name() { \ - auto out = std::make_shared(); \ - return out; \ + name out; \ + return std::dynamic_pointer_cast(out.impl()); \ } \ ops::OpPrimCRegisterHelper primc_gen_##name(#name, GetDefaultPrimC##name); } // namespace abstract diff --git a/mindspore/core/mindapi/base/types.h b/mindspore/core/mindapi/base/types.h index b42889b412..414aec8c8d 100644 --- a/mindspore/core/mindapi/base/types.h +++ b/mindspore/core/mindapi/base/types.h @@ -114,5 +114,10 @@ enum PaddingMode : int64_t { SYMMETRIC = 2, MODE_RESERVED = 3, }; + +enum PoolMode : int64_t { + MAX_POOLING = 0, + MEAN_POOLING = 1, +}; } // namespace mindspore #endif // MINDSPORE_CORE_MINDAPI_BASE_TYPES_H_ diff --git a/mindspore/core/mindapi/ir/abstract.h b/mindspore/core/mindapi/ir/abstract.h index 5d709e14e7..ddd6a4ecc7 100644 --- a/mindspore/core/mindapi/ir/abstract.h +++ b/mindspore/core/mindapi/ir/abstract.h @@ -153,5 +153,7 @@ class MIND_API AbstractTuple : public AbstractSequence { /// \param[in] elements A list of abstracts. explicit AbstractTuple(const AbstractBasePtrList &elements); }; + +using AbstractTuplePtr = SharedPtr; } // namespace mindspore::api #endif // MINDSPORE_CORE_MINDAPI_IR_ABSTRACT_H_ diff --git a/mindspore/core/mindapi/ir/common.h b/mindspore/core/mindapi/ir/common.h index c346c1654a..7dced654c0 100644 --- a/mindspore/core/mindapi/ir/common.h +++ b/mindspore/core/mindapi/ir/common.h @@ -47,5 +47,9 @@ using FuncGraphPtr = SharedPtr; class FuncGraphManager; using FuncGraphManagerPtr = SharedPtr; + +class CNode; +using CNodePtr = SharedPtr; +using CNodePtrList = std::vector; } // namespace mindspore::api #endif // MINDSPORE_CORE_MINDAPI_IR_COMMON_H_ diff --git a/mindspore/core/ops/LayerNormBetaGammaBackprop.cc b/mindspore/core/ops/LayerNormBetaGammaBackprop.cc index 903ecf0492..c08bd6f46c 100644 --- a/mindspore/core/ops/LayerNormBetaGammaBackprop.cc +++ b/mindspore/core/ops/LayerNormBetaGammaBackprop.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -49,6 +50,7 @@ TypePtr LayerNormBetaGammaBackpropInferType(const PrimitivePtr &prim, const std: } } // namespace +MIND_API_BASE_IMPL(LayerNormBetaGammaBackprop, PrimitiveC, BaseOperator); AbstractBasePtr LayerNormBetaGammaBackpropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/LayerNormBetaGammaBackprop.h b/mindspore/core/ops/LayerNormBetaGammaBackprop.h index ac9985219d..cd5f35bc00 100644 --- a/mindspore/core/ops/LayerNormBetaGammaBackprop.h +++ b/mindspore/core/ops/LayerNormBetaGammaBackprop.h @@ -20,23 +20,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API LayerNormBetaGammaBackprop : public PrimitiveC { +class MIND_API LayerNormBetaGammaBackprop : public BaseOperator { public: - LayerNormBetaGammaBackprop() : PrimitiveC(prim::kPrimLayerNormBetaGammaBackprop->name()) {} - ~LayerNormBetaGammaBackprop() = default; - MS_DECLARE_PARENT(LayerNormBetaGammaBackprop, PrimitiveC); + MIND_API_BASE_MEMBER(LayerNormBetaGammaBackprop); + LayerNormBetaGammaBackprop() : BaseOperator("LayerNormBetaGammaBackprop") {} void Init() const {} }; -AbstractBasePtr LayerNormBetaGammaBackpropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormBetaGammaBackpropInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/LayerNormXBackprop.cc b/mindspore/core/ops/LayerNormXBackprop.cc index 784d136a97..58fb4ed6a1 100644 --- a/mindspore/core/ops/LayerNormXBackprop.cc +++ b/mindspore/core/ops/LayerNormXBackprop.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,6 +45,7 @@ TypePtr LayerNormXBackpropInferType(const PrimitivePtr &prim, const std::vector< } } // namespace +MIND_API_BASE_IMPL(LayerNormXBackprop, PrimitiveC, BaseOperator); AbstractBasePtr LayerNormXBackpropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/LayerNormXBackprop.h b/mindspore/core/ops/LayerNormXBackprop.h index e50d7324bb..17a112bd12 100644 --- a/mindspore/core/ops/LayerNormXBackprop.h +++ b/mindspore/core/ops/LayerNormXBackprop.h @@ -20,23 +20,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API LayerNormXBackprop : public PrimitiveC { +class MIND_API LayerNormXBackprop : public BaseOperator { public: - LayerNormXBackprop() : PrimitiveC(prim::kPrimLayerNormXBackprop->name()) {} - ~LayerNormXBackprop() = default; - MS_DECLARE_PARENT(LayerNormXBackprop, PrimitiveC); + MIND_API_BASE_MEMBER(LayerNormXBackprop); + LayerNormXBackprop() : BaseOperator("LayerNormXBackprop") {} void Init() const {} }; -AbstractBasePtr LayerNormXBackpropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormXBackpropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/abs.cc b/mindspore/core/ops/abs.cc index 214eec05c4..577a9c7381 100644 --- a/mindspore/core/ops/abs.cc +++ b/mindspore/core/ops/abs.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -142,6 +143,8 @@ ValuePtr AbsInferValue(const PrimitivePtr &prim, const std::vector #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns absolute value of a tensor element-wise. /// Refer to Python API @ref mindspore.ops.Abs for more details. -class MS_CORE_API Abs : public PrimitiveC { +class MIND_API Abs : public BaseOperator { public: + MIND_API_BASE_MEMBER(Abs); /// \brief Constructor. - Abs() : PrimitiveC(prim::kPrimAbs->name()) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~Abs() = default; - MS_DECLARE_PARENT(Abs, PrimitiveC); + Abs() : BaseOperator("Abs") { InitIOName({"input_x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Abs for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/accumulate_n_v2.cc b/mindspore/core/ops/accumulate_n_v2.cc index b1eda5ec98..0ceabe3011 100644 --- a/mindspore/core/ops/accumulate_n_v2.cc +++ b/mindspore/core/ops/accumulate_n_v2.cc @@ -22,6 +22,8 @@ #include "ops/accumulate_n_v2.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -83,6 +85,7 @@ TypePtr AccumulateNV2InferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/accumulate_n_v2.h b/mindspore/core/ops/accumulate_n_v2.h index a3bdee0311..ec4ad4c3f0 100644 --- a/mindspore/core/ops/accumulate_n_v2.h +++ b/mindspore/core/ops/accumulate_n_v2.h @@ -19,21 +19,19 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAccumulateNV2 = "AccumulateNV2"; -class MS_CORE_API AccumulateNV2 : public PrimitiveC { +class MIND_API AccumulateNV2 : public BaseOperator { public: - AccumulateNV2() : PrimitiveC(kNameAccumulateNV2) { InitIOName({"inputs"}, {"sum"}); } - ~AccumulateNV2() = default; - MS_DECLARE_PARENT(AccumulateNV2, PrimitiveC); + MIND_API_BASE_MEMBER(AccumulateNV2); + AccumulateNV2() : BaseOperator(kNameAccumulateNV2) { InitIOName({"inputs"}, {"sum"}); } }; -AbstractBasePtr AccumulateNV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AccumulateNV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAccumulateNV2Ptr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/acos.cc b/mindspore/core/ops/acos.cc index 4b8efa7129..435289d6e0 100644 --- a/mindspore/core/ops/acos.cc +++ b/mindspore/core/ops/acos.cc @@ -15,6 +15,15 @@ */ #include "ops/acos.h" +#include +#include +#include +#include +#include +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -38,6 +47,7 @@ TypePtr ACosInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/acos.h b/mindspore/core/ops/acos.h index e2f221f740..45afd9bf75 100644 --- a/mindspore/core/ops/acos.h +++ b/mindspore/core/ops/acos.h @@ -22,28 +22,23 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameACos = "ACos"; /// \brief Computes arccosine of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.ACos for more details. -class ACos : public PrimitiveC { +class MIND_API ACos : public BaseOperator { public: + MIND_API_BASE_MEMBER(ACos); /// \brief Constructor. - ACos() : PrimitiveC(kNameACos) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~ACos() = default; - - MS_DECLARE_PARENT(ACos, PrimitiveC); + ACos() : BaseOperator(kNameACos) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr ACosInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ACosInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimACosPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/acosh.cc b/mindspore/core/ops/acosh.cc index b528e91e3c..7ea63e1174 100644 --- a/mindspore/core/ops/acosh.cc +++ b/mindspore/core/ops/acosh.cc @@ -15,6 +15,10 @@ */ #include "ops/acosh.h" +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +46,7 @@ TypePtr AcoshInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/acosh.h b/mindspore/core/ops/acosh.h index 00ca087a38..a8b3997877 100644 --- a/mindspore/core/ops/acosh.h +++ b/mindspore/core/ops/acosh.h @@ -22,28 +22,23 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAcosh = "Acosh"; /// \brief Computes arccosh of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.Acosh for more details. -class Acosh : public PrimitiveC { +class MIND_API Acosh : public BaseOperator { public: + MIND_API_BASE_MEMBER(Acosh); /// \brief Constructor. - Acosh() : PrimitiveC(kNameAcosh) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Acosh() = default; - - MS_DECLARE_PARENT(Acosh, PrimitiveC); + Acosh() : BaseOperator(kNameAcosh) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr AcoshInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AcoshInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAcoshPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/adam.cc b/mindspore/core/ops/adam.cc index c907e0b6a7..56c6d48b82 100644 --- a/mindspore/core/ops/adam.cc +++ b/mindspore/core/ops/adam.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -79,14 +80,18 @@ abstract::TupleShapePtr AdamInferShape(const PrimitivePtr &primitive, const std: std::vector{var_shape_ptr, m_shape_ptr, v_shape_ptr}); } } // namespace + +MIND_API_BASE_IMPL(Adam, PrimitiveC, BaseOperator); void Adam::Init(const bool use_locking, const bool use_nesterov) { this->set_use_locking(use_locking); this->set_use_nesterov(use_nesterov); } -void Adam::set_use_locking(const bool use_locking) { (void)this->AddAttr(kUseLocking, MakeValue(use_locking)); } +void Adam::set_use_locking(const bool use_locking) { (void)this->AddAttr(kUseLocking, api::MakeValue(use_locking)); } -void Adam::set_use_nesterov(const bool use_nesterov) { (void)this->AddAttr(kUseNesterov, MakeValue(use_nesterov)); } +void Adam::set_use_nesterov(const bool use_nesterov) { + (void)this->AddAttr(kUseNesterov, api::MakeValue(use_nesterov)); +} bool Adam::get_use_locking() const { auto value_ptr = GetAttr(kUseLocking); diff --git a/mindspore/core/ops/adam.h b/mindspore/core/ops/adam.h index 252a8765da..6065a3c36c 100644 --- a/mindspore/core/ops/adam.h +++ b/mindspore/core/ops/adam.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAdam = "Adam"; /// \brief Updates gradients by the Adaptive Moment Estimation (Adam) algorithm. /// Refer to Python API @ref mindspore.ops.Adam for more details. -class MS_CORE_API Adam : public PrimitiveC { +class MIND_API Adam : public BaseOperator { public: + MIND_API_BASE_MEMBER(Adam); /// \brief Constructor. - Adam() : PrimitiveC(kNameAdam) {} - /// \brief Destructor. - ~Adam() = default; - MS_DECLARE_PARENT(Adam, PrimitiveC); + Adam() : BaseOperator(kNameAdam) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Adam for the inputs. void Init(const bool use_locking = false, const bool use_nesterov = false); /// \brief Set use_locking. @@ -51,8 +48,8 @@ class MS_CORE_API Adam : public PrimitiveC { /// \return use_nesterov. bool get_use_nesterov() const; }; -AbstractBasePtr AdamInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AdamInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimAdamPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/add.cc b/mindspore/core/ops/add.cc index ff25fbb42c..62dac2c52e 100644 --- a/mindspore/core/ops/add.cc +++ b/mindspore/core/ops/add.cc @@ -21,9 +21,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Add, PrimitiveC, BaseOperator); AbstractBasePtr AddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/add.h b/mindspore/core/ops/add.h index df1540ae6b..286a2760b7 100644 --- a/mindspore/core/ops/add.h +++ b/mindspore/core/ops/add.h @@ -20,28 +20,25 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameAdd = prim::kAdd; +constexpr auto kNameAdd = "Add"; /// \brief Adds two input tensors element-wise. Refer to Python API @ref mindspore.ops.Add for more details. -class MS_CORE_API Add : public PrimitiveC { +class MIND_API Add : public BaseOperator { public: + MIND_API_BASE_MEMBER(Add); /// \brief Constructor. - Add() : PrimitiveC(kNameAdd) { InitIOName({"x", "y"}, {"output"}); } - explicit Add(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Add() = default; - MS_DECLARE_PARENT(Add, PrimitiveC); + Add() : BaseOperator(kNameAdd) { InitIOName({"x", "y"}, {"output"}); } + explicit Add(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Add for the inputs. void Init() const {} }; -AbstractBasePtr AddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/addcdiv.cc b/mindspore/core/ops/addcdiv.cc index d7155a0017..d68dd8bbe0 100644 --- a/mindspore/core/ops/addcdiv.cc +++ b/mindspore/core/ops/addcdiv.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -73,6 +74,8 @@ TypePtr AddcdivInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/addcdiv.h b/mindspore/core/ops/addcdiv.h index ff8b7195f3..5427f01962 100644 --- a/mindspore/core/ops/addcdiv.h +++ b/mindspore/core/ops/addcdiv.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAddcdiv = "Addcdiv"; -class Addcdiv : public PrimitiveC { +class MIND_API Addcdiv : public BaseOperator { public: - Addcdiv() : PrimitiveC(kNameAddcdiv) { InitIOName({"input_data", "x1", "x2", "value"}, {"output"}); } - ~Addcdiv() = default; - MS_DECLARE_PARENT(Addcdiv, PrimitiveC); + MIND_API_BASE_MEMBER(Addcdiv); + Addcdiv() : BaseOperator(kNameAddcdiv) { InitIOName({"input_data", "x1", "x2", "value"}, {"output"}); } }; -AbstractBasePtr AddcdivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AddcdivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAddcdivPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/addcmul.cc b/mindspore/core/ops/addcmul.cc index 05fb3a9190..ec0124a3a6 100644 --- a/mindspore/core/ops/addcmul.cc +++ b/mindspore/core/ops/addcmul.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -73,6 +74,8 @@ TypePtr AddcmulInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/addcmul.h b/mindspore/core/ops/addcmul.h index 6f252ab242..92a0644dcf 100644 --- a/mindspore/core/ops/addcmul.h +++ b/mindspore/core/ops/addcmul.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAddcmul = "Addcmul"; -class Addcmul : public PrimitiveC { +class MIND_API Addcmul : public BaseOperator { public: - Addcmul() : PrimitiveC(kNameAddcmul) { InitIOName({"input_data", "x1", "x2", "value"}, {"output"}); } - ~Addcmul() = default; - MS_DECLARE_PARENT(Addcmul, PrimitiveC); + MIND_API_BASE_MEMBER(Addcmul); + Addcmul() : BaseOperator(kNameAddcmul) { InitIOName({"input_data", "x1", "x2", "value"}, {"output"}); } }; -AbstractBasePtr AddcmulInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AddcmulInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAddcmulPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/adder.cc b/mindspore/core/ops/adder.cc index 9285713c45..215ca3f228 100644 --- a/mindspore/core/ops/adder.cc +++ b/mindspore/core/ops/adder.cc @@ -16,9 +16,11 @@ #include "ops/adder.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Adder, PrimitiveC, BaseOperator); void Adder::Init(const int64_t in_channel, const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &stride, const std::vector &pad_list, const std::vector &dilation, const int64_t group, const Format &format) { @@ -33,14 +35,16 @@ void Adder::Init(const int64_t in_channel, const int64_t out_channel, const std: set_format(format); } -void Adder::set_in_channel(const int64_t in_channel) { (void)this->AddAttr(kInChannel, MakeValue(in_channel)); } +void Adder::set_in_channel(const int64_t in_channel) { (void)this->AddAttr(kInChannel, api::MakeValue(in_channel)); } int64_t Adder::get_in_channel() const { auto value_ptr = GetAttr(kInChannel); return GetValue(value_ptr); } -void Adder::set_out_channel(const int64_t out_channel) { (void)this->AddAttr(kOutChannel, MakeValue(out_channel)); } +void Adder::set_out_channel(const int64_t out_channel) { + (void)this->AddAttr(kOutChannel, api::MakeValue(out_channel)); +} int64_t Adder::get_out_channel() const { auto value_ptr = GetAttr(kOutChannel); @@ -48,7 +52,7 @@ int64_t Adder::get_out_channel() const { } void Adder::set_kernel_size(const std::vector &kernel_size) { - (void)this->AddAttr(kKernelSize, MakeValue(kernel_size)); + (void)this->AddAttr(kKernelSize, api::MakeValue(kernel_size)); } std::vector Adder::get_kernel_size() const { @@ -58,7 +62,7 @@ std::vector Adder::get_kernel_size() const { void Adder::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode Adder::get_pad_mode() const { @@ -66,28 +70,32 @@ PadMode Adder::get_pad_mode() const { return PadMode(GetValue(value_ptr)); } -void Adder::set_stride(const std::vector &stride) { (void)this->AddAttr(kStride, MakeValue(stride)); } +void Adder::set_stride(const std::vector &stride) { (void)this->AddAttr(kStride, api::MakeValue(stride)); } std::vector Adder::get_stride() const { auto value_ptr = GetAttr(kStride); return GetValue>(value_ptr); } -void Adder::set_pad_list(const std::vector &pad_list) { (void)this->AddAttr(kPadList, MakeValue(pad_list)); } +void Adder::set_pad_list(const std::vector &pad_list) { + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); +} std::vector Adder::get_pad_list() const { auto value_ptr = GetAttr(kPadList); return GetValue>(value_ptr); } -void Adder::set_dilation(const std::vector &dilation) { (void)this->AddAttr(kDilation, MakeValue(dilation)); } +void Adder::set_dilation(const std::vector &dilation) { + (void)this->AddAttr(kDilation, api::MakeValue(dilation)); +} std::vector Adder::get_dilation() const { auto value_ptr = GetAttr(kDilation); return GetValue>(value_ptr); } -void Adder::set_group(const int64_t group) { (void)this->AddAttr(kGroup, MakeValue(group)); } +void Adder::set_group(const int64_t group) { (void)this->AddAttr(kGroup, api::MakeValue(group)); } int64_t Adder::get_group() const { auto value_ptr = GetAttr(kGroup); @@ -96,7 +104,7 @@ int64_t Adder::get_group() const { void Adder::set_format(const Format &format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } Format Adder::get_format() const { diff --git a/mindspore/core/ops/adder.h b/mindspore/core/ops/adder.h index c0faddd20b..88b166af4e 100644 --- a/mindspore/core/ops/adder.h +++ b/mindspore/core/ops/adder.h @@ -21,22 +21,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameAdder = "Adder"; /// \brief All defined All operator prototype of lite. -class MS_CORE_API Adder : public PrimitiveC { +class MIND_API Adder : public BaseOperator { public: + MIND_API_BASE_MEMBER(Adder); /// \brief Constructor. - explicit Adder(const std::string &k_name = kNameAdder) : PrimitiveC(k_name) {} - - /// \brief Destructor. - ~Adder() = default; - MS_DECLARE_PARENT(Adder, PrimitiveC); + explicit Adder(const std::string &k_name = kNameAdder) : BaseOperator(k_name) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/addn.cc b/mindspore/core/ops/addn.cc index 8251ccda1a..21aca84e2e 100644 --- a/mindspore/core/ops/addn.cc +++ b/mindspore/core/ops/addn.cc @@ -21,6 +21,8 @@ #include #include "ops/addn.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -83,6 +85,8 @@ TypePtr AddNInferType(const PrimitivePtr &prim, const std::vectorBuildType(); } } // namespace + +MIND_API_BASE_IMPL(AddN, PrimitiveC, BaseOperator); AbstractBasePtr AddNInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/addn.h b/mindspore/core/ops/addn.h index 3bcfc95e1d..4c6600dea7 100644 --- a/mindspore/core/ops/addn.h +++ b/mindspore/core/ops/addn.h @@ -18,27 +18,24 @@ #define MINDSPORE_CORE_OPS_ADDN_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAddN = "AddN"; /// \brief Computes addition of all input tensors element-wise. /// Refer to Python API @ref mindspore.ops.AddN for more details. -class MS_CORE_API AddN : public PrimitiveC { +class MIND_API AddN : public BaseOperator { public: + MIND_API_BASE_MEMBER(AddN); /// \brief Constructor. - AddN() : PrimitiveC(kNameAddN) { InitIOName({"inputs"}, {"sum"}); } - /// \brief Destructor. - ~AddN() = default; - MS_DECLARE_PARENT(AddN, PrimitiveC); + AddN() : BaseOperator(kNameAddN) { InitIOName({"inputs"}, {"sum"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.AddN for the inputs. void Init() const {} }; -AbstractBasePtr AddNInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AddNInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/affine.cc b/mindspore/core/ops/affine.cc index dd0bf463a0..524c5c0f99 100644 --- a/mindspore/core/ops/affine.cc +++ b/mindspore/core/ops/affine.cc @@ -17,8 +17,11 @@ #include "ops/affine.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Affine, PrimitiveC, BaseOperator); void Affine::Init(const std::vector &contexts, int64_t output_dim, bool transpose_a, bool transpose_b) { this->set_context(contexts); this->set_output_dim(output_dim); @@ -27,17 +30,17 @@ void Affine::Init(const std::vector &contexts, int64_t output_dim, bool } void Affine::set_context(const std::vector &context) { - (void)this->AddAttr(kAffineContext, MakeValue(context)); + (void)this->AddAttr(kAffineContext, api::MakeValue(context)); } -void Affine::set_output_dim(int64_t output_dim) { (void)this->AddAttr(kAffineOutputDim, MakeValue(output_dim)); } +void Affine::set_output_dim(int64_t output_dim) { (void)this->AddAttr(kAffineOutputDim, api::MakeValue(output_dim)); } -void Affine::set_transpose_a(bool transpose_a) { (void)AddAttr(kTransposeA, MakeValue(transpose_a)); } +void Affine::set_transpose_a(bool transpose_a) { (void)AddAttr(kTransposeA, api::MakeValue(transpose_a)); } -void Affine::set_transpose_b(bool transpose_b) { (void)AddAttr(kTransposeB, MakeValue(transpose_b)); } +void Affine::set_transpose_b(bool transpose_b) { (void)AddAttr(kTransposeB, api::MakeValue(transpose_b)); } void Affine::set_activation_type(const ActivationType &activation_type) { - (void)this->AddAttr(kActivationType, MakeValue(static_cast(activation_type))); + (void)this->AddAttr(kActivationType, api::MakeValue(static_cast(activation_type))); } bool Affine::get_transpose_a() const { diff --git a/mindspore/core/ops/affine.h b/mindspore/core/ops/affine.h index 22ec7aeb01..366d07f16b 100644 --- a/mindspore/core/ops/affine.h +++ b/mindspore/core/ops/affine.h @@ -18,25 +18,21 @@ #define MINDSPORE_CORE_OPS_AFFINE_H_ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { - constexpr auto kNameAffine = "Affine"; constexpr auto kAffineContext = "context"; constexpr auto kAffineOutputDim = "output_dim"; /// \brief Assert defined Affine operator prototype of lite. -class MS_CORE_API Affine : public PrimitiveC { +class MIND_API Affine : public BaseOperator { public: + MIND_API_BASE_MEMBER(Affine); /// \brief Constructor. - Affine() : PrimitiveC(kNameAffine) { InitIOName({"x1", "x2"}, {"outputs"}); } - /// \brief Destructor. - ~Affine() = default; - MS_DECLARE_PARENT(Affine, PrimitiveC); + Affine() : BaseOperator(kNameAffine) { InitIOName({"x1", "x2"}, {"outputs"}); } /// \brief Method to init the op's attributes. void Init(const std::vector &contexts, int64_t output_dim, bool transpose_a = false, bool transpose_b = false); diff --git a/mindspore/core/ops/all.cc b/mindspore/core/ops/all.cc index 702438701e..30a6d999c7 100644 --- a/mindspore/core/ops/all.cc +++ b/mindspore/core/ops/all.cc @@ -17,12 +17,14 @@ #include "ops/all.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(All, PrimitiveC, BaseOperator); void All::Init(const int64_t keep_dims) { this->set_keep_dims(keep_dims); } -void All::set_keep_dims(const int64_t keep_dims) { (void)this->AddAttr(kKeepDims, MakeValue(keep_dims)); } +void All::set_keep_dims(const int64_t keep_dims) { (void)this->AddAttr(kKeepDims, api::MakeValue(keep_dims)); } int64_t All::get_keep_dims() const { auto value_ptr = GetAttr(kKeepDims); diff --git a/mindspore/core/ops/all.h b/mindspore/core/ops/all.h index 671e544743..295f1cee8b 100644 --- a/mindspore/core/ops/all.h +++ b/mindspore/core/ops/all.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_ALL_H_ #define MINDSPORE_CORE_OPS_ALL_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAll = "All"; /// \brief All defined All operator prototype of lite. -class MS_CORE_API All : public PrimitiveC { +class MIND_API All : public BaseOperator { public: + MIND_API_BASE_MEMBER(All); /// \brief Constructor. - All() : PrimitiveC(kNameAll) {} - - /// \brief Destructor. - ~All() = default; - - MS_DECLARE_PARENT(All, PrimitiveC); + All() : BaseOperator(kNameAll) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/all_gather.cc b/mindspore/core/ops/all_gather.cc index 1da487988a..9ee5712dfc 100644 --- a/mindspore/core/ops/all_gather.cc +++ b/mindspore/core/ops/all_gather.cc @@ -17,12 +17,14 @@ #include "ops/all_gather.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(AllGather, PrimitiveC, BaseOperator); void AllGather::set_group(const string &group) { std::string g = group; - (void)this->AddAttr(kGroup, MakeValue(g)); + (void)this->AddAttr(kGroup, api::MakeValue(g)); } std::string AllGather::get_group() const { auto value_ptr = GetAttr(kGroup); @@ -30,7 +32,7 @@ std::string AllGather::get_group() const { } void AllGather::set_rank_size(int rank_size) { - (void)this->AddAttr(kRankSize, MakeValue(static_cast(rank_size))); + (void)this->AddAttr(kRankSize, api::MakeValue(static_cast(rank_size))); } int AllGather::get_rank_size() const { auto value_ptr = GetAttr(kRankSize); diff --git a/mindspore/core/ops/all_gather.h b/mindspore/core/ops/all_gather.h index abb1c5213c..7f9052cf83 100644 --- a/mindspore/core/ops/all_gather.h +++ b/mindspore/core/ops/all_gather.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAllGather = "AllGather"; -class MS_CORE_API AllGather : public PrimitiveC { +class MIND_API AllGather : public BaseOperator { public: - AllGather() : PrimitiveC(kNameAllGather) { InitIOName({"input_x"}, {"output"}); } - ~AllGather() = default; - MS_DECLARE_PARENT(AllGather, PrimitiveC); + MIND_API_BASE_MEMBER(AllGather); + AllGather() : BaseOperator(kNameAllGather) { InitIOName({"input_x"}, {"output"}); } void Init() {} void set_group(const std::string &format); std::string get_group() const; diff --git a/mindspore/core/ops/apply_ada_max.cc b/mindspore/core/ops/apply_ada_max.cc index b1eae72f3d..fc13893bb5 100644 --- a/mindspore/core/ops/apply_ada_max.cc +++ b/mindspore/core/ops/apply_ada_max.cc @@ -1,164 +1,167 @@ -/** - * 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 "ops/apply_ada_max.h" - -#include -#include -#include "abstract/primitive_infer_map.h" -#include "ops/op_utils.h" -#include "utils/tensor_construct_utils.h" - -namespace mindspore { -namespace ops { -namespace { -abstract::TupleShapePtr ApplyAdaMaxInferShape(const PrimitivePtr &primitive, - const std::vector &input_args) { - MS_EXCEPTION_IF_NULL(primitive); - const int64_t kInputNum = 9; - (void)CheckAndConvertUtils::CheckInteger("input number", SizeToLong(input_args.size()), kGreaterEqual, kInputNum, - primitive->name()); - for (const auto &item : input_args) { - MS_EXCEPTION_IF_NULL(item); - } - auto prim_name = primitive->name(); - auto var_shape = input_args[kInputIndex0]->BuildShape(); - auto m_shape = input_args[kInputIndex1]->BuildShape(); - auto v_shape = input_args[kInputIndex2]->BuildShape(); - auto var_shape_ptr = var_shape->cast(); - auto m_shape_ptr = m_shape->cast(); - auto v_shape_ptr = v_shape->cast(); - auto beta1_power_shape = - CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex3]->BuildShape())[kShape]; - auto lr_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex4]->BuildShape())[kShape]; - auto beta1_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex5]->BuildShape())[kShape]; - auto beta2_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex6]->BuildShape())[kShape]; - auto epsilon_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex7]->BuildShape())[kShape]; - auto grad_shape = input_args[kInputIndex8]->BuildShape(); - auto grad_shape_ptr = grad_shape->cast(); - // beta1_power,lr,beta1,beta2,epsilon must be scalar - const int64_t kInputShape = 1; - (void)CheckAndConvertUtils::CheckInteger("beta1 power's rank", beta1_power_shape.size(), kLessEqual, kInputShape, - prim_name); - if (beta1_power_shape.size() == 1) { - (void)CheckAndConvertUtils::CheckInteger("beta1_power_shape[0]", beta1_power_shape.size(), kEqual, kInputShape, - prim_name); - } - (void)CheckAndConvertUtils::CheckInteger("lr's rank", lr_shape.size(), kLessEqual, kInputShape, prim_name); - if (lr_shape.size() == 1) { - (void)CheckAndConvertUtils::CheckInteger("lr_shape[0]", lr_shape.size(), kEqual, kInputShape, prim_name); - } - (void)CheckAndConvertUtils::CheckInteger("beta1's rank", beta1_shape.size(), kLessEqual, kInputShape, prim_name); - if (beta1_shape.size() == 1) { - (void)CheckAndConvertUtils::CheckInteger("beta1_shape[0]", beta1_shape.size(), kEqual, kInputShape, prim_name); - } - (void)CheckAndConvertUtils::CheckInteger("beta2's rank", beta2_shape.size(), kLessEqual, kInputShape, prim_name); - if (beta2_shape.size() == 1) { - (void)CheckAndConvertUtils::CheckInteger("beta2_shape[0]", beta2_shape.size(), kEqual, kInputShape, prim_name); - } - (void)CheckAndConvertUtils::CheckInteger("epsilon's rank", epsilon_shape.size(), kLessEqual, kInputShape, prim_name); - if (epsilon_shape.size() == 1) { - (void)CheckAndConvertUtils::CheckInteger("epsilon_shape[0]", epsilon_shape.size(), kEqual, kInputShape, prim_name); - } - - // var, m,v and grad must have the same shape - std::map same_shape_args_map; - same_shape_args_map.insert({"m", m_shape}); - same_shape_args_map.insert({"v", v_shape}); - same_shape_args_map.insert({"grad", grad_shape}); - if (!var_shape_ptr->IsDynamic() && !m_shape_ptr->IsDynamic()) { - if (*m_shape != *var_shape) { - MS_EXCEPTION(ValueError) << primitive->name() << " evaluator arg m shape " << m_shape->ToString() - << " are not consistent with var shape " << var_shape->ToString(); - } - } - if (!v_shape_ptr->IsDynamic() && !var_shape_ptr->IsDynamic()) { - if (*v_shape != *var_shape) { - MS_EXCEPTION(ValueError) << primitive->name() << " evaluator arg v shape " << v_shape->ToString() - << " are not consistent with var shape " << var_shape->ToString(); - } - } - if (!grad_shape_ptr->IsDynamic() && !var_shape_ptr->IsDynamic()) { - if (*grad_shape != *var_shape) { - MS_EXCEPTION(ValueError) << primitive->name() << " evaluator arg grad shape " << grad_shape->ToString() - << " are not consistent with var shape " << var_shape->ToString(); - } - } - - return std::make_shared(std::vector{var_shape, m_shape, v_shape}); -} - -TuplePtr ApplyAdaMaxInferType(const PrimitivePtr &prim, const std::vector &input_args) { - MS_EXCEPTION_IF_NULL(prim); - auto prim_name = prim->name(); - const int64_t kInputNum = 9; - (void)CheckAndConvertUtils::CheckInteger("input number", SizeToLong(input_args.size()), kGreaterEqual, kInputNum, - prim_name); - for (const auto &item : input_args) { - MS_EXCEPTION_IF_NULL(item); - } - auto var_type = input_args[kInputIndex0]->BuildType(); - auto m_type = input_args[kInputIndex1]->BuildType(); - auto v_type = input_args[kInputIndex2]->BuildType(); - auto beta1_power_type = input_args[kInputIndex3]->BuildType(); - auto lr_type = input_args[kInputIndex4]->BuildType(); - auto beta1_type = input_args[kInputIndex5]->BuildType(); - auto beta2_type = input_args[kInputIndex6]->BuildType(); - auto epsilon_type = input_args[kInputIndex7]->BuildType(); - auto grad_type = input_args[kInputIndex8]->BuildType(); - const std::set valid_types = {kFloat16, kFloat32}; - // m v grad must have the same type as var - std::map args; - (void)args.insert({"var_type", var_type}); - (void)args.insert({"m_type", m_type}); - (void)args.insert({"v_type", v_type}); - (void)args.insert({"grad_type", grad_type}); - (void)CheckAndConvertUtils::CheckTensorTypeSame(args, valid_types, prim_name); - - std::map args_beta1_power; - std::map args_lr; - std::map args_beta1; - std::map args_beta2; - std::map args_epsilon; - - (void)args_beta1_power.insert({"beta1_power_type", beta1_power_type}); - (void)args_lr.insert({"lr_type", lr_type}); - (void)args_beta1.insert({"beta1_type", beta1_type}); - (void)args_beta2.insert({"beta2_type", beta2_type}); - (void)args_epsilon.insert({"epsilon_type", epsilon_type}); - - // beta1_power,lr,beta1,beta2,epsilon must be a scalar or zero dimension tensor type - (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_beta1_power, valid_types, prim_name); - (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_lr, valid_types, prim_name); - (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_beta1, valid_types, prim_name); - (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_beta2, valid_types, prim_name); - (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_epsilon, valid_types, prim_name); - - return std::make_shared(std::vector{var_type, m_type, v_type}); -} -} // namespace - -AbstractBasePtr ApplyAdaMaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args) { - MS_EXCEPTION_IF_NULL(primitive); - auto infer_type = ApplyAdaMaxInferType(primitive, input_args); - auto infer_shape = ApplyAdaMaxInferShape(primitive, input_args); - return abstract::MakeAbstract(infer_shape, infer_type); -} - -REGISTER_PRIMITIVE_EVAL_IMPL(ApplyAdaMax, prim::kPrimApplyAdaMax, ApplyAdaMaxInfer, nullptr, true); -} // namespace ops -} // namespace mindspore +/** + * 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 "ops/apply_ada_max.h" + +#include +#include +#include "abstract/primitive_infer_map.h" +#include "ops/op_utils.h" +#include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" + +namespace mindspore { +namespace ops { +namespace { +abstract::TupleShapePtr ApplyAdaMaxInferShape(const PrimitivePtr &primitive, + const std::vector &input_args) { + MS_EXCEPTION_IF_NULL(primitive); + const int64_t kInputNum = 9; + (void)CheckAndConvertUtils::CheckInteger("input number", SizeToLong(input_args.size()), kGreaterEqual, kInputNum, + primitive->name()); + for (const auto &item : input_args) { + MS_EXCEPTION_IF_NULL(item); + } + auto prim_name = primitive->name(); + auto var_shape = input_args[kInputIndex0]->BuildShape(); + auto m_shape = input_args[kInputIndex1]->BuildShape(); + auto v_shape = input_args[kInputIndex2]->BuildShape(); + auto var_shape_ptr = var_shape->cast(); + auto m_shape_ptr = m_shape->cast(); + auto v_shape_ptr = v_shape->cast(); + auto beta1_power_shape = + CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex3]->BuildShape())[kShape]; + auto lr_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex4]->BuildShape())[kShape]; + auto beta1_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex5]->BuildShape())[kShape]; + auto beta2_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex6]->BuildShape())[kShape]; + auto epsilon_shape = CheckAndConvertUtils::ConvertShapePtrToShapeMap(input_args[kInputIndex7]->BuildShape())[kShape]; + auto grad_shape = input_args[kInputIndex8]->BuildShape(); + auto grad_shape_ptr = grad_shape->cast(); + // beta1_power,lr,beta1,beta2,epsilon must be scalar + const int64_t kInputShape = 1; + (void)CheckAndConvertUtils::CheckInteger("beta1 power's rank", beta1_power_shape.size(), kLessEqual, kInputShape, + prim_name); + if (beta1_power_shape.size() == 1) { + (void)CheckAndConvertUtils::CheckInteger("beta1_power_shape[0]", beta1_power_shape.size(), kEqual, kInputShape, + prim_name); + } + (void)CheckAndConvertUtils::CheckInteger("lr's rank", lr_shape.size(), kLessEqual, kInputShape, prim_name); + if (lr_shape.size() == 1) { + (void)CheckAndConvertUtils::CheckInteger("lr_shape[0]", lr_shape.size(), kEqual, kInputShape, prim_name); + } + (void)CheckAndConvertUtils::CheckInteger("beta1's rank", beta1_shape.size(), kLessEqual, kInputShape, prim_name); + if (beta1_shape.size() == 1) { + (void)CheckAndConvertUtils::CheckInteger("beta1_shape[0]", beta1_shape.size(), kEqual, kInputShape, prim_name); + } + (void)CheckAndConvertUtils::CheckInteger("beta2's rank", beta2_shape.size(), kLessEqual, kInputShape, prim_name); + if (beta2_shape.size() == 1) { + (void)CheckAndConvertUtils::CheckInteger("beta2_shape[0]", beta2_shape.size(), kEqual, kInputShape, prim_name); + } + (void)CheckAndConvertUtils::CheckInteger("epsilon's rank", epsilon_shape.size(), kLessEqual, kInputShape, prim_name); + if (epsilon_shape.size() == 1) { + (void)CheckAndConvertUtils::CheckInteger("epsilon_shape[0]", epsilon_shape.size(), kEqual, kInputShape, prim_name); + } + + // var, m,v and grad must have the same shape + std::map same_shape_args_map; + same_shape_args_map.insert({"m", m_shape}); + same_shape_args_map.insert({"v", v_shape}); + same_shape_args_map.insert({"grad", grad_shape}); + if (!var_shape_ptr->IsDynamic() && !m_shape_ptr->IsDynamic()) { + if (*m_shape != *var_shape) { + MS_EXCEPTION(ValueError) << primitive->name() << " evaluator arg m shape " << m_shape->ToString() + << " are not consistent with var shape " << var_shape->ToString(); + } + } + if (!v_shape_ptr->IsDynamic() && !var_shape_ptr->IsDynamic()) { + if (*v_shape != *var_shape) { + MS_EXCEPTION(ValueError) << primitive->name() << " evaluator arg v shape " << v_shape->ToString() + << " are not consistent with var shape " << var_shape->ToString(); + } + } + if (!grad_shape_ptr->IsDynamic() && !var_shape_ptr->IsDynamic()) { + if (*grad_shape != *var_shape) { + MS_EXCEPTION(ValueError) << primitive->name() << " evaluator arg grad shape " << grad_shape->ToString() + << " are not consistent with var shape " << var_shape->ToString(); + } + } + + return std::make_shared(std::vector{var_shape, m_shape, v_shape}); +} + +TuplePtr ApplyAdaMaxInferType(const PrimitivePtr &prim, const std::vector &input_args) { + MS_EXCEPTION_IF_NULL(prim); + auto prim_name = prim->name(); + const int64_t kInputNum = 9; + (void)CheckAndConvertUtils::CheckInteger("input number", SizeToLong(input_args.size()), kGreaterEqual, kInputNum, + prim_name); + for (const auto &item : input_args) { + MS_EXCEPTION_IF_NULL(item); + } + auto var_type = input_args[kInputIndex0]->BuildType(); + auto m_type = input_args[kInputIndex1]->BuildType(); + auto v_type = input_args[kInputIndex2]->BuildType(); + auto beta1_power_type = input_args[kInputIndex3]->BuildType(); + auto lr_type = input_args[kInputIndex4]->BuildType(); + auto beta1_type = input_args[kInputIndex5]->BuildType(); + auto beta2_type = input_args[kInputIndex6]->BuildType(); + auto epsilon_type = input_args[kInputIndex7]->BuildType(); + auto grad_type = input_args[kInputIndex8]->BuildType(); + const std::set valid_types = {kFloat16, kFloat32}; + // m v grad must have the same type as var + std::map args; + (void)args.insert({"var_type", var_type}); + (void)args.insert({"m_type", m_type}); + (void)args.insert({"v_type", v_type}); + (void)args.insert({"grad_type", grad_type}); + (void)CheckAndConvertUtils::CheckTensorTypeSame(args, valid_types, prim_name); + + std::map args_beta1_power; + std::map args_lr; + std::map args_beta1; + std::map args_beta2; + std::map args_epsilon; + + (void)args_beta1_power.insert({"beta1_power_type", beta1_power_type}); + (void)args_lr.insert({"lr_type", lr_type}); + (void)args_beta1.insert({"beta1_type", beta1_type}); + (void)args_beta2.insert({"beta2_type", beta2_type}); + (void)args_epsilon.insert({"epsilon_type", epsilon_type}); + + // beta1_power,lr,beta1,beta2,epsilon must be a scalar or zero dimension tensor type + (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_beta1_power, valid_types, prim_name); + (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_lr, valid_types, prim_name); + (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_beta1, valid_types, prim_name); + (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_beta2, valid_types, prim_name); + (void)CheckAndConvertUtils::CheckScalarOrTensorTypesSame(args_epsilon, valid_types, prim_name); + + return std::make_shared(std::vector{var_type, m_type, v_type}); +} +} // namespace + +MIND_API_BASE_IMPL(ApplyAdaMax, PrimitiveC, BaseOperator); +AbstractBasePtr ApplyAdaMaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args) { + MS_EXCEPTION_IF_NULL(primitive); + auto infer_type = ApplyAdaMaxInferType(primitive, input_args); + auto infer_shape = ApplyAdaMaxInferShape(primitive, input_args); + return abstract::MakeAbstract(infer_shape, infer_type); +} + +REGISTER_PRIMITIVE_EVAL_IMPL(ApplyAdaMax, prim::kPrimApplyAdaMax, ApplyAdaMaxInfer, nullptr, true); +} // namespace ops +} // namespace mindspore diff --git a/mindspore/core/ops/apply_ada_max.h b/mindspore/core/ops/apply_ada_max.h index 69d2527817..ee9d62fe3e 100644 --- a/mindspore/core/ops/apply_ada_max.h +++ b/mindspore/core/ops/apply_ada_max.h @@ -1,45 +1,43 @@ -/** - * 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. - */ - -#ifndef MINDSPORE_CORE_OPS_APPLY_ADA_MAX_H_ -#define MINDSPORE_CORE_OPS_APPLY_ADA_MAX_H_ -#include -#include -#include -#include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" - -namespace mindspore { -namespace ops { -constexpr auto kNameApplyAdaMax = "ApplyAdaMax"; -class ApplyAdaMax : public PrimitiveC { - public: - ApplyAdaMax() : PrimitiveC(kNameApplyAdaMax) { - InitIOName({"var", "m", "v", "beta1_power", "lr", "beta1", "beta2", "epsilon", "grad"}, {"var", "m", "v"}); - } - ~ApplyAdaMax() = default; - MS_DECLARE_PARENT(ApplyAdaMax, PrimitiveC); -}; -AbstractBasePtr ApplyAdaMaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); - -using kPrimApplyAdaMaxPtr = std::shared_ptr; -} // namespace ops -} // namespace mindspore - -#endif // MINDSPORE_CORE_OPS_APPLY_ADA_MAX_H_ +/** + * 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. + */ + +#ifndef MINDSPORE_CORE_OPS_APPLY_ADA_MAX_H_ +#define MINDSPORE_CORE_OPS_APPLY_ADA_MAX_H_ +#include +#include +#include +#include +#include "ops/base_operator.h" +#include "mindapi/base/types.h" + +namespace mindspore { +namespace ops { +constexpr auto kNameApplyAdaMax = "ApplyAdaMax"; +class MIND_API ApplyAdaMax : public BaseOperator { + public: + MIND_API_BASE_MEMBER(ApplyAdaMax); + ApplyAdaMax() : BaseOperator(kNameApplyAdaMax) { + InitIOName({"var", "m", "v", "beta1_power", "lr", "beta1", "beta2", "epsilon", "grad"}, {"var", "m", "v"}); + } +}; +abstract::AbstractBasePtr ApplyAdaMaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); + +using kPrimApplyAdaMaxPtr = std::shared_ptr; +} // namespace ops +} // namespace mindspore + +#endif // MINDSPORE_CORE_OPS_APPLY_ADA_MAX_H_ diff --git a/mindspore/core/ops/apply_adadelta.cc b/mindspore/core/ops/apply_adadelta.cc index 26c8b00e8c..2adf960c82 100644 --- a/mindspore/core/ops/apply_adadelta.cc +++ b/mindspore/core/ops/apply_adadelta.cc @@ -23,6 +23,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -121,6 +122,8 @@ TuplePtr ApplyAdadeltaInferType(const PrimitivePtr &primitive, const std::vector return std::make_shared(std::vector{var_type, accum_type, accum_update_type}); } } // namespace + +MIND_API_BASE_IMPL(ApplyAdadelta, PrimitiveC, BaseOperator); AbstractBasePtr ApplyAdadeltaInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = ApplyAdadeltaInferType(primitive, input_args); diff --git a/mindspore/core/ops/apply_adadelta.h b/mindspore/core/ops/apply_adadelta.h index 3d9fab9034..ac19836281 100644 --- a/mindspore/core/ops/apply_adadelta.h +++ b/mindspore/core/ops/apply_adadelta.h @@ -22,23 +22,21 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyAdadelta = "ApplyAdadelta"; -class ApplyAdadelta : public PrimitiveC { +class MIND_API ApplyAdadelta : public BaseOperator { public: - ApplyAdadelta() : PrimitiveC(kNameApplyAdadelta) { + MIND_API_BASE_MEMBER(ApplyAdadelta); + ApplyAdadelta() : BaseOperator(kNameApplyAdadelta) { InitIOName({"var", "accum", "accum_update", "lr", "rho", "epsilon", "grad"}, {"var", "accum", "accum_update"}); } - ~ApplyAdadelta() = default; - MS_DECLARE_PARENT(ApplyAdadelta, PrimitiveC); }; -AbstractBasePtr ApplyAdadeltaInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyAdadeltaInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyAdadeltaPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_adagrad.cc b/mindspore/core/ops/apply_adagrad.cc index 03b381be1d..9d2552dda2 100644 --- a/mindspore/core/ops/apply_adagrad.cc +++ b/mindspore/core/ops/apply_adagrad.cc @@ -22,6 +22,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -81,6 +83,7 @@ TuplePtr ApplyAdagradInferType(const PrimitivePtr &primitive, const std::vector< } } // namespace +MIND_API_BASE_IMPL(ApplyAdagrad, PrimitiveC, BaseOperator); AbstractBasePtr ApplyAdagradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_adagrad.h b/mindspore/core/ops/apply_adagrad.h index 9457322539..e67aa08e4a 100644 --- a/mindspore/core/ops/apply_adagrad.h +++ b/mindspore/core/ops/apply_adagrad.h @@ -22,22 +22,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyAdagrad = "ApplyAdagrad"; -class ApplyAdagrad : public PrimitiveC { +class MIND_API ApplyAdagrad : public BaseOperator { public: - ApplyAdagrad() : PrimitiveC(kNameApplyAdagrad) { InitIOName({"var", "accum", "lr", "grad"}, {"var", "accum"}); } - ~ApplyAdagrad() = default; - MS_DECLARE_PARENT(ApplyAdagrad, PrimitiveC); + MIND_API_BASE_MEMBER(ApplyAdagrad); + ApplyAdagrad() : BaseOperator(kNameApplyAdagrad) { InitIOName({"var", "accum", "lr", "grad"}, {"var", "accum"}); } }; -AbstractBasePtr ApplyAdagradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyAdagradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyAdagradPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/apply_adagrad_d_a.cc b/mindspore/core/ops/apply_adagrad_d_a.cc index bf2219657f..956764db78 100644 --- a/mindspore/core/ops/apply_adagrad_d_a.cc +++ b/mindspore/core/ops/apply_adagrad_d_a.cc @@ -23,10 +23,10 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { - namespace { abstract::TupleShapePtr ApplyAdagradDAInferShape(const PrimitivePtr &primitive, const std::vector &input_args) { @@ -98,6 +98,7 @@ TuplePtr ApplyAdagradDAInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_adagrad_d_a.h b/mindspore/core/ops/apply_adagrad_d_a.h index 7031a15086..1ff15db7f2 100644 --- a/mindspore/core/ops/apply_adagrad_d_a.h +++ b/mindspore/core/ops/apply_adagrad_d_a.h @@ -22,31 +22,26 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyAdagradDA = "ApplyAdagradDA"; /// \brief Update var according to the proximal adagrad scheme. /// Refer to Python API @ref mindspore.ops.ApplyAdagradDA for more details. -class ApplyAdagradDA : public PrimitiveC { +class MIND_API ApplyAdagradDA : public BaseOperator { public: + MIND_API_BASE_MEMBER(ApplyAdagradDA); /// \brief Constructor. - ApplyAdagradDA() : PrimitiveC(kNameApplyAdagradDA) { + ApplyAdagradDA() : BaseOperator(kNameApplyAdagradDA) { InitIOName({"var", "gradient_accumulator", "gradient_squared_accumulator", "grad", "lr", "l1", "l2", "global_step"}, {"var", "gradient_accumulator", "gradient_squared_accumulator"}); } - - /// \brief Destructor. - ~ApplyAdagradDA() = default; - - MS_DECLARE_PARENT(ApplyAdagradDA, PrimitiveC); }; -AbstractBasePtr ApplyAdagradDAInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyAdagradDAInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_adagrad_v2.cc b/mindspore/core/ops/apply_adagrad_v2.cc index b146ec59d7..2695e8a744 100644 --- a/mindspore/core/ops/apply_adagrad_v2.cc +++ b/mindspore/core/ops/apply_adagrad_v2.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -77,6 +78,7 @@ TuplePtr ApplyAdagradV2InferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_adagrad_v2.h b/mindspore/core/ops/apply_adagrad_v2.h index d544d3538e..7aa969082f 100644 --- a/mindspore/core/ops/apply_adagrad_v2.h +++ b/mindspore/core/ops/apply_adagrad_v2.h @@ -22,23 +22,19 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyAdagradV2 = "ApplyAdagradV2"; -class ApplyAdagradV2 : public PrimitiveC { +class MIND_API ApplyAdagradV2 : public BaseOperator { public: - ApplyAdagradV2() : PrimitiveC(kNameApplyAdagradV2) { InitIOName({"var", "accum", "lr", "grad"}, {"var", "accum"}); } - - ~ApplyAdagradV2() = default; - - MS_DECLARE_PARENT(ApplyAdagradV2, PrimitiveC); + MIND_API_BASE_MEMBER(ApplyAdagradV2); + ApplyAdagradV2() : BaseOperator(kNameApplyAdagradV2) { InitIOName({"var", "accum", "lr", "grad"}, {"var", "accum"}); } }; -AbstractBasePtr ApplyAdagradV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyAdagradV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyAdagradV2Ptr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_adam_with_amsgrad.cc b/mindspore/core/ops/apply_adam_with_amsgrad.cc index df1bf7e845..a33bc049ad 100644 --- a/mindspore/core/ops/apply_adam_with_amsgrad.cc +++ b/mindspore/core/ops/apply_adam_with_amsgrad.cc @@ -23,6 +23,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -91,6 +93,7 @@ TuplePtr ApplyAdamWithAmsgradInferType(const PrimitivePtr &prim, const std::vect } } // namespace +MIND_API_BASE_IMPL(ApplyAdamWithAmsgrad, PrimitiveC, BaseOperator); AbstractBasePtr ApplyAdamWithAmsgradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_adam_with_amsgrad.h b/mindspore/core/ops/apply_adam_with_amsgrad.h index 705d3253bd..d54840e42a 100644 --- a/mindspore/core/ops/apply_adam_with_amsgrad.h +++ b/mindspore/core/ops/apply_adam_with_amsgrad.h @@ -19,24 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyAdamWithAmsgrad = "ApplyAdamWithAmsgrad"; -class ApplyAdamWithAmsgrad : public PrimitiveC { +class MIND_API ApplyAdamWithAmsgrad : public BaseOperator { public: - ApplyAdamWithAmsgrad() : PrimitiveC(kNameApplyAdamWithAmsgrad) { + MIND_API_BASE_MEMBER(ApplyAdamWithAmsgrad); + ApplyAdamWithAmsgrad() : BaseOperator(kNameApplyAdamWithAmsgrad) { InitIOName({"var", "m", "v", "vhat", "beta1_power", "beta2_power", "lr", "grad"}, {"var", "m", "v", "vhat"}); } - ~ApplyAdamWithAmsgrad() = default; - MS_DECLARE_PARENT(ApplyAdamWithAmsgrad, PrimitiveC); }; -AbstractBasePtr ApplyAdamWithAmsgradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyAdamWithAmsgradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimApplyAdamWithAmsgradPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/apply_add_sign.cc b/mindspore/core/ops/apply_add_sign.cc index 346e6de0e7..2ae140c6a8 100644 --- a/mindspore/core/ops/apply_add_sign.cc +++ b/mindspore/core/ops/apply_add_sign.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -116,6 +117,7 @@ TuplePtr ApplyAddSignInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_add_sign.h b/mindspore/core/ops/apply_add_sign.h index 6a936f2b3d..eac2f69d99 100644 --- a/mindspore/core/ops/apply_add_sign.h +++ b/mindspore/core/ops/apply_add_sign.h @@ -21,27 +21,23 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyAddSign = "ApplyAddSign"; -class ApplyAddSign : public PrimitiveC { +class MIND_API ApplyAddSign : public BaseOperator { public: - ApplyAddSign() : PrimitiveC(kNameApplyAddSign) { + MIND_API_BASE_MEMBER(ApplyAddSign); + ApplyAddSign() : BaseOperator(kNameApplyAddSign) { InitIOName({"var", "m", "lr", "alpha", "sign_decay", "beta", "grad"}, {"var", "m"}); } - - ~ApplyAddSign() = default; - - MS_DECLARE_PARENT(ApplyAddSign, PrimitiveC); }; -AbstractBasePtr ApplyAddSignInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyAddSignInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyAddSignPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_centered_rms_prop.cc b/mindspore/core/ops/apply_centered_rms_prop.cc index 4aee8cc698..d037b902ef 100644 --- a/mindspore/core/ops/apply_centered_rms_prop.cc +++ b/mindspore/core/ops/apply_centered_rms_prop.cc @@ -17,6 +17,7 @@ #include "ops/apply_centered_rms_prop.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -106,6 +107,8 @@ TypePtr ApplyCenteredRMSPropInferType(const PrimitivePtr &primitive, const std:: return var_dtype; } } // namespace + +MIND_API_BASE_IMPL(ApplyCenteredRMSProp, PrimitiveC, BaseOperator); AbstractBasePtr ApplyCenteredRMSPropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = ApplyCenteredRMSPropInferType(primitive, input_args); diff --git a/mindspore/core/ops/apply_centered_rms_prop.h b/mindspore/core/ops/apply_centered_rms_prop.h index 20d4c1f9c1..3b1d02c04f 100644 --- a/mindspore/core/ops/apply_centered_rms_prop.h +++ b/mindspore/core/ops/apply_centered_rms_prop.h @@ -22,25 +22,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyCenteredRMSProp = "ApplyCenteredRMSProp"; -class ApplyCenteredRMSProp : public PrimitiveC { +class MIND_API ApplyCenteredRMSProp : public BaseOperator { public: - ApplyCenteredRMSProp() : PrimitiveC(kNameApplyCenteredRMSProp) { + MIND_API_BASE_MEMBER(ApplyCenteredRMSProp); + ApplyCenteredRMSProp() : BaseOperator(kNameApplyCenteredRMSProp) { InitIOName( {"var", "mean_gradient", "mean_square", "moment", "grad", "learning_rate", "decay", "momentum", "epsilon"}, {"var", "mean_gradient", "mean_square", "moment"}); } - ~ApplyCenteredRMSProp() = default; - MS_DECLARE_PARENT(ApplyCenteredRMSProp, PrimitiveC); }; -AbstractBasePtr ApplyCenteredRMSPropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyCenteredRMSPropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyCenteredRMSPropPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_ftrl.cc b/mindspore/core/ops/apply_ftrl.cc index dccdb2e634..04ddc9c31e 100644 --- a/mindspore/core/ops/apply_ftrl.cc +++ b/mindspore/core/ops/apply_ftrl.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -91,6 +92,8 @@ TypePtr ApplyFtrlInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_ftrl.h b/mindspore/core/ops/apply_ftrl.h index 7317667577..287549586e 100644 --- a/mindspore/core/ops/apply_ftrl.h +++ b/mindspore/core/ops/apply_ftrl.h @@ -22,24 +22,21 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyFtrl = "ApplyFtrl"; -class ApplyFtrl : public PrimitiveC { +class MIND_API ApplyFtrl : public BaseOperator { public: - ApplyFtrl() : PrimitiveC(kNameApplyFtrl) { + MIND_API_BASE_MEMBER(ApplyFtrl); + ApplyFtrl() : BaseOperator(kNameApplyFtrl) { InitIOName({"var", "accum", "linear", "grad", "lr", "l1", "l2", "lr_power"}, {"var"}); } - - ~ApplyFtrl() = default; - MS_DECLARE_PARENT(ApplyFtrl, PrimitiveC); }; -AbstractBasePtr ApplyFtrlInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyFtrlInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyFtrlPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_gradient_descent.cc b/mindspore/core/ops/apply_gradient_descent.cc index cddc95580a..bc7bcb361e 100644 --- a/mindspore/core/ops/apply_gradient_descent.cc +++ b/mindspore/core/ops/apply_gradient_descent.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -72,6 +73,7 @@ TypePtr ApplyGradientDescentInferType(const PrimitivePtr &prim, const std::vecto } } // namespace +MIND_API_BASE_IMPL(ApplyGradientDescent, PrimitiveC, BaseOperator); AbstractBasePtr ApplyGradientDescentInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_gradient_descent.h b/mindspore/core/ops/apply_gradient_descent.h index ad45cbbb85..a0b0f9b533 100644 --- a/mindspore/core/ops/apply_gradient_descent.h +++ b/mindspore/core/ops/apply_gradient_descent.h @@ -22,24 +22,20 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyGradientDescent = "ApplyGradientDescent"; -class ApplyGradientDescent : public PrimitiveC { +class MIND_API ApplyGradientDescent : public BaseOperator { public: - ApplyGradientDescent() : PrimitiveC(kNameApplyGradientDescent) { InitIOName({"var", "alpha", "delta"}, {"var"}); } - - ~ApplyGradientDescent() = default; - - MS_DECLARE_PARENT(ApplyGradientDescent, PrimitiveC); + MIND_API_BASE_MEMBER(ApplyGradientDescent); + ApplyGradientDescent() : BaseOperator(kNameApplyGradientDescent) { InitIOName({"var", "alpha", "delta"}, {"var"}); } }; -AbstractBasePtr ApplyGradientDescentInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyGradientDescentInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimApplyGradientDescentPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/apply_keras_momentum.cc b/mindspore/core/ops/apply_keras_momentum.cc index 4f6c6679f1..a99c322da9 100644 --- a/mindspore/core/ops/apply_keras_momentum.cc +++ b/mindspore/core/ops/apply_keras_momentum.cc @@ -22,6 +22,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -81,6 +83,7 @@ TuplePtr ApplyKerasMomentumInferType(const PrimitivePtr &prim, const std::vector } } // namespace +MIND_API_BASE_IMPL(ApplyKerasMomentum, PrimitiveC, BaseOperator); AbstractBasePtr ApplyKerasMomentumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_keras_momentum.h b/mindspore/core/ops/apply_keras_momentum.h index 75d4c74df4..02d81c8ef1 100644 --- a/mindspore/core/ops/apply_keras_momentum.h +++ b/mindspore/core/ops/apply_keras_momentum.h @@ -22,24 +22,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyKerasMomentum = "ApplyKerasMomentum"; -class MS_CORE_API ApplyKerasMomentum : public PrimitiveC { +class MIND_API ApplyKerasMomentum : public BaseOperator { public: - ApplyKerasMomentum() : PrimitiveC(kNameApplyKerasMomentum) { + MIND_API_BASE_MEMBER(ApplyKerasMomentum); + ApplyKerasMomentum() : BaseOperator(kNameApplyKerasMomentum) { InitIOName({"var", "accum", "lr", "grad", "momentum"}, {"var", "accum"}); } - ~ApplyKerasMomentum() = default; - MS_DECLARE_PARENT(ApplyKerasMomentum, PrimitiveC); }; -AbstractBasePtr ApplyKerasMomentumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyKerasMomentumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimApplyKerasMomentumPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/apply_momentum.cc b/mindspore/core/ops/apply_momentum.cc index cb2aab62d2..d70b7bfb2c 100644 --- a/mindspore/core/ops/apply_momentum.cc +++ b/mindspore/core/ops/apply_momentum.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -30,15 +31,15 @@ void ApplyMomentum::Init(const bool use_nesterov, const bool use_locking, const } void ApplyMomentum::set_use_nesterov(const bool use_nesterov) { - (void)this->AddAttr(kUseNesterov, MakeValue(use_nesterov)); + (void)this->AddAttr(kUseNesterov, api::MakeValue(use_nesterov)); } void ApplyMomentum::set_use_locking(const bool use_locking) { - (void)this->AddAttr(kUseLocking, MakeValue(use_locking)); + (void)this->AddAttr(kUseLocking, api::MakeValue(use_locking)); } void ApplyMomentum::set_gradient_scale(const float gradient_scale) { - (void)this->AddAttr(kGradientScale, MakeValue(gradient_scale)); + (void)this->AddAttr(kGradientScale, api::MakeValue(gradient_scale)); } bool ApplyMomentum::get_use_nesterov() const { @@ -102,6 +103,8 @@ TypePtr ApplyMomentumInferType(const PrimitivePtr &primitive, const std::vector< return v_tensor_type; } } // namespace + +MIND_API_BASE_IMPL(ApplyMomentum, PrimitiveC, BaseOperator); AbstractBasePtr ApplyMomentumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = ApplyMomentumInferType(primitive, input_args); diff --git a/mindspore/core/ops/apply_momentum.h b/mindspore/core/ops/apply_momentum.h index 7dbb3bb372..199afa1c81 100644 --- a/mindspore/core/ops/apply_momentum.h +++ b/mindspore/core/ops/apply_momentum.h @@ -22,24 +22,21 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyMomentum = "ApplyMomentum"; /// \brief Optimizer that implements the Momentum algorithm. /// Refer to Python API @ref mindspore.ops.ApplyMomentum for more details. -class MS_CORE_API ApplyMomentum : public PrimitiveC { +class MIND_API ApplyMomentum : public BaseOperator { public: + MIND_API_BASE_MEMBER(ApplyMomentum); /// \brief Constructor. - ApplyMomentum() : PrimitiveC(kNameApplyMomentum) { + ApplyMomentum() : BaseOperator(kNameApplyMomentum) { InitIOName({"var", "accum", "lr", "grad", "momentum"}, {"var", "accum"}); } - /// \brief Destructor. - ~ApplyMomentum() = default; - MS_DECLARE_PARENT(ApplyMomentum, PrimitiveC); /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ApplyMomentum for the inputs. void Init(const bool use_nesterov = false, const bool use_locking = false, const float gradient_scale = 1.0); /// \brief Set use_nesterov. @@ -61,8 +58,8 @@ class MS_CORE_API ApplyMomentum : public PrimitiveC { /// \return gradient_scale. float get_gradient_scale() const; }; -AbstractBasePtr ApplyMomentumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyMomentumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyMomentumPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_power_sign_d.cc b/mindspore/core/ops/apply_power_sign_d.cc index 6872b8ffd6..aa795dcb70 100644 --- a/mindspore/core/ops/apply_power_sign_d.cc +++ b/mindspore/core/ops/apply_power_sign_d.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -108,6 +109,7 @@ TuplePtr ApplyPowerSignDInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_power_sign_d.h b/mindspore/core/ops/apply_power_sign_d.h index a06eb4d4aa..c01310304a 100644 --- a/mindspore/core/ops/apply_power_sign_d.h +++ b/mindspore/core/ops/apply_power_sign_d.h @@ -19,23 +19,21 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyPowerSign = "ApplyPowerSign"; -class ApplyPowerSign : public PrimitiveC { +class MIND_API ApplyPowerSign : public BaseOperator { public: - ApplyPowerSign() : PrimitiveC(kNameApplyPowerSign) { + MIND_API_BASE_MEMBER(ApplyPowerSign); + ApplyPowerSign() : BaseOperator(kNameApplyPowerSign) { InitIOName({"var", "m", "lr", "logbase", "sign_decay", "beta", "grad"}, {"var", "m"}); } - ~ApplyPowerSign() = default; - MS_DECLARE_PARENT(ApplyPowerSign, PrimitiveC); }; -AbstractBasePtr ApplyPowerSignDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyPowerSignDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyPowerSignDPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/apply_proximal_adagrad.cc b/mindspore/core/ops/apply_proximal_adagrad.cc index e3bf786ca7..56eac981c1 100644 --- a/mindspore/core/ops/apply_proximal_adagrad.cc +++ b/mindspore/core/ops/apply_proximal_adagrad.cc @@ -22,6 +22,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -99,6 +101,7 @@ TuplePtr ApplyProximalAdagradInferType(const PrimitivePtr &primitive, const std: } } // namespace +MIND_API_BASE_IMPL(ApplyProximalAdagrad, PrimitiveC, BaseOperator); AbstractBasePtr ApplyProximalAdagradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/apply_proximal_adagrad.h b/mindspore/core/ops/apply_proximal_adagrad.h index d8936c4161..22f5ead2b8 100644 --- a/mindspore/core/ops/apply_proximal_adagrad.h +++ b/mindspore/core/ops/apply_proximal_adagrad.h @@ -22,24 +22,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyProximalAdagrad = "ApplyProximalAdagrad"; -class ApplyProximalAdagrad : public PrimitiveC { +class MIND_API ApplyProximalAdagrad : public BaseOperator { public: - ApplyProximalAdagrad() : PrimitiveC(kNameApplyProximalAdagrad) { + MIND_API_BASE_MEMBER(ApplyProximalAdagrad); + ApplyProximalAdagrad() : BaseOperator(kNameApplyProximalAdagrad) { InitIOName({"var", "accum", "lr", "l1", "l2", "grad"}, {"var", "accum"}); } - ~ApplyProximalAdagrad() = default; - MS_DECLARE_PARENT(ApplyProximalAdagrad, PrimitiveC); }; -AbstractBasePtr ApplyProximalAdagradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyProximalAdagradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApplyProximalAdagradPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/apply_proximal_gradient_descent.cc b/mindspore/core/ops/apply_proximal_gradient_descent.cc index 55dc4315fe..eb286a611a 100644 --- a/mindspore/core/ops/apply_proximal_gradient_descent.cc +++ b/mindspore/core/ops/apply_proximal_gradient_descent.cc @@ -25,6 +25,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -98,6 +99,7 @@ TypePtr ApplyProximalGradientDescentInferType(const PrimitivePtr &prim, } } // namespace +MIND_API_BASE_IMPL(ApplyProximalGradientDescent, PrimitiveC, BaseOperator); AbstractBasePtr ApplyProximalGradientDescentInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { const int64_t input_num = 5; diff --git a/mindspore/core/ops/apply_proximal_gradient_descent.h b/mindspore/core/ops/apply_proximal_gradient_descent.h index 802938b37b..e32205f98e 100644 --- a/mindspore/core/ops/apply_proximal_gradient_descent.h +++ b/mindspore/core/ops/apply_proximal_gradient_descent.h @@ -19,23 +19,22 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameApplyProximalGradientDescent = "ApplyProximalGradientDescent"; -class ApplyProximalGradientDescent : public PrimitiveC { +class MIND_API ApplyProximalGradientDescent : public BaseOperator { public: - ApplyProximalGradientDescent() : PrimitiveC(kNameApplyProximalGradientDescent) { + MIND_API_BASE_MEMBER(ApplyProximalGradientDescent); + ApplyProximalGradientDescent() : BaseOperator(kNameApplyProximalGradientDescent) { InitIOName({"var", "alpha", "l1", "l2", "delta"}, {"var"}); } - ~ApplyProximalGradientDescent() = default; - MS_DECLARE_PARENT(ApplyProximalGradientDescent, PrimitiveC); }; -AbstractBasePtr ApplyProximalGradientDescentInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApplyProximalGradientDescentInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/approximate_equal.cc b/mindspore/core/ops/approximate_equal.cc index 3d550722ee..23f4b7f0bb 100644 --- a/mindspore/core/ops/approximate_equal.cc +++ b/mindspore/core/ops/approximate_equal.cc @@ -18,6 +18,9 @@ #include #include #include +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -59,6 +62,8 @@ TypePtr ApproximateEqualInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/approximate_equal.h b/mindspore/core/ops/approximate_equal.h index de3a76cb21..d35227899f 100644 --- a/mindspore/core/ops/approximate_equal.h +++ b/mindspore/core/ops/approximate_equal.h @@ -19,21 +19,18 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { -class ApproximateEqual : public PrimitiveC { +class MIND_API ApproximateEqual : public BaseOperator { public: - ApproximateEqual() : PrimitiveC(prim::kPrimApproximateEqual->name()) {} - ~ApproximateEqual() = default; - MS_DECLARE_PARENT(ApproximateEqual, PrimitiveC); + MIND_API_BASE_MEMBER(ApproximateEqual); + ApproximateEqual() : BaseOperator("ApproximateEqual") {} void Init() {} }; -AbstractBasePtr ApproximateEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ApproximateEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimApproximateEqualPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/arg_max.cc b/mindspore/core/ops/arg_max.cc index 110a4ba40f..a450d99517 100644 --- a/mindspore/core/ops/arg_max.cc +++ b/mindspore/core/ops/arg_max.cc @@ -15,6 +15,10 @@ */ #include "ops/arg_max.h" +#include "mindapi/ir/type.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -23,15 +27,17 @@ void ArgMax::Init(const int64_t axis, const TypeId output_type) { set_output_type(output_type); } -void ArgMax::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } -void ArgMax::set_output_type(const TypeId output_type) { (void)this->AddAttr(kOutputType, TypeIdToType(output_type)); } +void ArgMax::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } +void ArgMax::set_output_type(const TypeId output_type) { + (void)this->AddAttr(kOutputType, api::Type::GetType(output_type)); +} int64_t ArgMax::get_axis() const { return GetValue(GetAttr(kAxis)); } TypeId ArgMax::get_output_type() const { - auto type_ptr = GetAttr(kOutputType)->cast()->element(); + auto type_ptr = GetAttr(kOutputType)->cast()->element(); return type_ptr->type_id(); } - +MIND_API_BASE_IMPL(ArgMax, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameArgMax, ArgMax); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/arg_max.h b/mindspore/core/ops/arg_max.h index 7a3a2bd0b2..8a7b90a3c2 100644 --- a/mindspore/core/ops/arg_max.h +++ b/mindspore/core/ops/arg_max.h @@ -20,24 +20,21 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/type_id.h" namespace mindspore { namespace ops { constexpr auto kNameArgMax = "Argmax"; /// \brief Returns the indices of the maximum value of a tensor across the axis. /// Refer to Python API @ref mindspore.ops.Argmax for more details. -class MS_CORE_API ArgMax : public PrimitiveC { +class MIND_API ArgMax : public BaseOperator { public: + MIND_API_BASE_MEMBER(ArgMax); /// \brief Constructor. - ArgMax() : PrimitiveC(kNameArgMax) { InitIOName({"x"}, {"output"}); } - explicit ArgMax(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~ArgMax() = default; - MS_DECLARE_PARENT(ArgMax, PrimitiveC); + ArgMax() : BaseOperator(kNameArgMax) { InitIOName({"x"}, {"output"}); } + explicit ArgMax(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Argmax for the inputs. void Init(const int64_t axis = -1, const TypeId output_type = kNumberTypeInt32); /// \brief Set axis. @@ -54,8 +51,8 @@ class MS_CORE_API ArgMax : public PrimitiveC { /// \return output_type. TypeId get_output_type() const; }; -AbstractBasePtr ArgMaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ArgMaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/arg_min.cc b/mindspore/core/ops/arg_min.cc index 1b97e51feb..71b5fc47b6 100644 --- a/mindspore/core/ops/arg_min.cc +++ b/mindspore/core/ops/arg_min.cc @@ -16,21 +16,28 @@ #include #include "ops/arg_min.h" +#include "mindapi/ir/type.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ArgMin, PrimitiveC, BaseOperator); void ArgMin::Init(const int64_t axis, const TypeId output_type) { set_axis(axis); set_output_type(output_type); } -void ArgMin::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } -void ArgMin::set_output_type(const TypeId output_type) { (void)this->AddAttr(kOutputType, TypeIdToType(output_type)); } +void ArgMin::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } +void ArgMin::set_output_type(const TypeId output_type) { + (void)this->AddAttr(kOutputType, api::Type::GetType(output_type)); +} int64_t ArgMin::get_axis() const { return GetValue(GetAttr(kAxis)); } TypeId ArgMin::get_output_type() const { - auto type_ptr = GetAttr(kOutputType)->cast()->element(); + auto type_ptr = GetAttr(kOutputType)->cast()->element(); return type_ptr->type_id(); } diff --git a/mindspore/core/ops/arg_min.h b/mindspore/core/ops/arg_min.h index 0e9fc09303..7be30fc6c2 100644 --- a/mindspore/core/ops/arg_min.h +++ b/mindspore/core/ops/arg_min.h @@ -20,24 +20,21 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/type_id.h" namespace mindspore { namespace ops { constexpr auto kNameArgMin = "ArgMin"; /// \brief Returns the indices of the minimum value of a tensor across the axis. /// Refer to Python API @ref mindspore.ops.Argmin for more details. -class MS_CORE_API ArgMin : public PrimitiveC { +class MIND_API ArgMin : public BaseOperator { public: + MIND_API_BASE_MEMBER(ArgMin); /// \brief Constructor. - ArgMin() : PrimitiveC(kNameArgMin) { InitIOName({"x"}, {"output"}); } - explicit ArgMin(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~ArgMin() = default; - MS_DECLARE_PARENT(ArgMin, PrimitiveC); + ArgMin() : BaseOperator(kNameArgMin) { InitIOName({"x"}, {"output"}); } + explicit ArgMin(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Argmin for the inputs. void Init(const int64_t axis = -1, const TypeId output_type = kNumberTypeInt32); /// \brief Set axis. @@ -54,8 +51,8 @@ class MS_CORE_API ArgMin : public PrimitiveC { /// \return output_type. TypeId get_output_type() const; }; -AbstractBasePtr ArgMinInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ArgMinInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimArgMin = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/asin.cc b/mindspore/core/ops/asin.cc index 74aceb9195..500008414e 100644 --- a/mindspore/core/ops/asin.cc +++ b/mindspore/core/ops/asin.cc @@ -19,6 +19,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "abstract/param_validator.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +43,7 @@ TypePtr AsinInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/asin.h b/mindspore/core/ops/asin.h index 7f67e8124b..6d31d59f6d 100644 --- a/mindspore/core/ops/asin.h +++ b/mindspore/core/ops/asin.h @@ -22,30 +22,26 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAsin = "Asin"; /// \brief Computes arcsine of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.Asin for more details. -class MS_CORE_API Asin : public PrimitiveC { +class MIND_API Asin : public BaseOperator { public: + MIND_API_BASE_MEMBER(Asin); /// \brief Constructor. - Asin() : PrimitiveC(kNameAsin) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Asin() = default; - - MS_DECLARE_PARENT(Asin, PrimitiveC); + Asin() : BaseOperator(kNameAsin) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Asin for the inputs. void Init() const {} }; -AbstractBasePtr AsinInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AsinInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAsinPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/asinh.cc b/mindspore/core/ops/asinh.cc index 170e872691..da3e61ef89 100644 --- a/mindspore/core/ops/asinh.cc +++ b/mindspore/core/ops/asinh.cc @@ -15,6 +15,11 @@ */ #include "ops/asinh.h" +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "abstract/param_validator.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +47,7 @@ TypePtr AsinhInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/asinh.h b/mindspore/core/ops/asinh.h index a7a77ac371..40e2574d11 100644 --- a/mindspore/core/ops/asinh.h +++ b/mindspore/core/ops/asinh.h @@ -22,29 +22,24 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAsinh = "Asinh"; /// \brief Computes arcsinh of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.Asinh for more details. -class MS_CORE_API Asinh : public PrimitiveC { +class MIND_API Asinh : public BaseOperator { public: - /// \brief Constructor. - Asinh() : PrimitiveC(kNameAsinh) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Asinh() = default; - - MS_DECLARE_PARENT(Asinh, PrimitiveC); + MIND_API_BASE_MEMBER(Asinh); + Asinh() : BaseOperator(kNameAsinh) { InitIOName({"x"}, {"y"}); } void Init() {} }; -AbstractBasePtr AsinhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AsinhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAsinhPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/assert.cc b/mindspore/core/ops/assert.cc index f65247e664..74d374ccd0 100644 --- a/mindspore/core/ops/assert.cc +++ b/mindspore/core/ops/assert.cc @@ -21,13 +21,16 @@ #include #include "ops/assert.h" +#include "mindapi/src/helper.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Assert, PrimitiveC, BaseOperator); void Assert::Init(const int64_t summarize) { set_summarize(summarize); } -void Assert::set_summarize(const int64_t summarize) { (void)this->AddAttr(kSummarize, MakeValue(summarize)); } +void Assert::set_summarize(const int64_t summarize) { (void)this->AddAttr(kSummarize, api::MakeValue(summarize)); } int64_t Assert::get_summarize() const { auto value_ptr = GetAttr(kSummarize); diff --git a/mindspore/core/ops/assert.h b/mindspore/core/ops/assert.h index bb421dce4c..f15bbc7465 100644 --- a/mindspore/core/ops/assert.h +++ b/mindspore/core/ops/assert.h @@ -18,23 +18,18 @@ #define MINDSPORE_CORE_OPS_ASSERT_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAssert = "Assert"; /// \brief Assert defined Assert operator prototype of lite. -class MS_CORE_API Assert : public PrimitiveC { +class MIND_API Assert : public BaseOperator { public: + MIND_API_BASE_MEMBER(Assert); /// \brief Constructor. - Assert() : PrimitiveC(kNameAssert) {} - - /// \brief Destructor. - ~Assert() = default; - - MS_DECLARE_PARENT(Assert, PrimitiveC); + Assert() : BaseOperator(kNameAssert) {} /// \brief Method to init the op's attributes. /// @@ -52,8 +47,8 @@ class MS_CORE_API Assert : public PrimitiveC { int64_t get_summarize() const; }; -AbstractBasePtr AssertInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AssertInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/assign.cc b/mindspore/core/ops/assign.cc index 9e34b7e8ad..7cd6dad8b0 100644 --- a/mindspore/core/ops/assign.cc +++ b/mindspore/core/ops/assign.cc @@ -23,9 +23,12 @@ #include "ops/assign.h" #include "ops/op_utils.h" #include "ir/dtype/ref.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Assign, PrimitiveC, BaseOperator); abstract::ShapePtr AssignInferShape(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(prim); auto prim_name = prim->name(); diff --git a/mindspore/core/ops/assign.h b/mindspore/core/ops/assign.h index c369c6a5a3..b900c02968 100644 --- a/mindspore/core/ops/assign.h +++ b/mindspore/core/ops/assign.h @@ -19,21 +19,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAssign = "Assign"; /// \brief Assigns Parameter with a value. Refer to Python API @ref mindspore.ops.Assign for more details. -class MS_CORE_API Assign : public PrimitiveC { +class MIND_API Assign : public BaseOperator { public: + MIND_API_BASE_MEMBER(Assign); /// \brief Constructor. - Assign() : PrimitiveC(kNameAssign) { InitIOName({"ref", "value"}, {"output"}); } - /// \brief Destructor. - ~Assign() = default; - MS_DECLARE_PARENT(Assign, PrimitiveC); + Assign() : BaseOperator(kNameAssign) { InitIOName({"ref", "value"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Assign for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/assign_add.cc b/mindspore/core/ops/assign_add.cc index 04e46a6149..7e75fa97d0 100644 --- a/mindspore/core/ops/assign_add.cc +++ b/mindspore/core/ops/assign_add.cc @@ -19,6 +19,7 @@ #include "ops/assign_add.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -39,6 +40,8 @@ TypePtr AssignAddInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/assign_add.h b/mindspore/core/ops/assign_add.h index 9fa49b5f58..4454b7451b 100644 --- a/mindspore/core/ops/assign_add.h +++ b/mindspore/core/ops/assign_add.h @@ -19,27 +19,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAssignAdd = "AssignAdd"; /// \brief Updates a Parameter by adding a value to it. /// Refer to Python API @ref mindspore.ops.AssignAdd for more details. -class MS_CORE_API AssignAdd : public PrimitiveC { +class MIND_API AssignAdd : public BaseOperator { public: + MIND_API_BASE_MEMBER(AssignAdd); /// \brief Constructor. - AssignAdd() : PrimitiveC(kNameAssignAdd) { InitIOName({"ref", "value"}, {"output"}); } - /// \brief Destructor. - ~AssignAdd() = default; - MS_DECLARE_PARENT(AssignAdd, PrimitiveC); + AssignAdd() : BaseOperator(kNameAssignAdd) { InitIOName({"ref", "value"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.AssignAdd for the inputs. void Init() const {} }; -AbstractBasePtr AssignAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AssignAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimAssignAddPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/assign_sub.cc b/mindspore/core/ops/assign_sub.cc index 96e07771fa..6a174aeb60 100644 --- a/mindspore/core/ops/assign_sub.cc +++ b/mindspore/core/ops/assign_sub.cc @@ -19,6 +19,7 @@ #include #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +42,7 @@ TypePtr AssignSubInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/assign_sub.h b/mindspore/core/ops/assign_sub.h index cb2b104c8d..a43f600924 100644 --- a/mindspore/core/ops/assign_sub.h +++ b/mindspore/core/ops/assign_sub.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAssignSub = "AssignSub"; -class AssignSub : public PrimitiveC { +class MIND_API AssignSub : public BaseOperator { public: - AssignSub() : PrimitiveC(kNameAssignSub) { InitIOName({"val", "value"}, {"val"}); } - ~AssignSub() = default; - MS_DECLARE_PARENT(AssignSub, PrimitiveC); + MIND_API_BASE_MEMBER(AssignSub); + AssignSub() : BaseOperator(kNameAssignSub) { InitIOName({"val", "value"}, {"val"}); } void Init() {} }; -AbstractBasePtr AssignSubInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AssignSubInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimAssignSubPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/atan.cc b/mindspore/core/ops/atan.cc index 368e7f1130..e97e057362 100644 --- a/mindspore/core/ops/atan.cc +++ b/mindspore/core/ops/atan.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -52,6 +53,8 @@ TypePtr AtanInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/atan.h b/mindspore/core/ops/atan.h index 4a0d6b244a..b1d3dfd0d3 100644 --- a/mindspore/core/ops/atan.h +++ b/mindspore/core/ops/atan.h @@ -20,27 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAtan = "Atan"; /// \brief Computes the trigonometric inverse tangent of the input element-wise. /// Refer to Python API @ref mindspore.ops.Atan for more details. -class MS_CORE_API Atan : public PrimitiveC { +class MIND_API Atan : public BaseOperator { public: + MIND_API_BASE_MEMBER(Atan); /// \brief Constructor. - Atan() : PrimitiveC(kNameAtan) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Atan() = default; - MS_DECLARE_PARENT(Atan, PrimitiveC); + Atan() : BaseOperator(kNameAtan) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Atan for the inputs. void Init() const {} }; -AbstractBasePtr AtanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AtanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/atanh.cc b/mindspore/core/ops/atanh.cc index 4af756b28f..e06af48149 100644 --- a/mindspore/core/ops/atanh.cc +++ b/mindspore/core/ops/atanh.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -58,6 +59,8 @@ TypePtr AtanhInferType(const PrimitivePtr &primitive, const std::vector &input_args) { auto type = AtanhInferType(primitive, input_args); diff --git a/mindspore/core/ops/atanh.h b/mindspore/core/ops/atanh.h index bbd5548919..bd7288ad4c 100644 --- a/mindspore/core/ops/atanh.h +++ b/mindspore/core/ops/atanh.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAtanh = "Atanh"; -class Atanh : public PrimitiveC { +class MIND_API Atanh : public BaseOperator { public: - Atanh() : PrimitiveC(kNameAtanh) { InitIOName({"x"}, {"output"}); } - ~Atanh() = default; - MS_DECLARE_PARENT(Atanh, PrimitiveC); + MIND_API_BASE_MEMBER(Atanh); + Atanh() : BaseOperator(kNameAtanh) { InitIOName({"x"}, {"output"}); } void Init() {} }; -AbstractBasePtr AtanhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AtanhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAtanhPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/attention.cc b/mindspore/core/ops/attention.cc index 89533ba520..df92084d7e 100644 --- a/mindspore/core/ops/attention.cc +++ b/mindspore/core/ops/attention.cc @@ -16,7 +16,10 @@ */ #include "ops/attention.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore::ops { +MIND_API_BASE_IMPL(Attention, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameAttention, Attention); } // namespace mindspore::ops diff --git a/mindspore/core/ops/attention.h b/mindspore/core/ops/attention.h index e86dc9d2c6..fbb7221202 100644 --- a/mindspore/core/ops/attention.h +++ b/mindspore/core/ops/attention.h @@ -19,25 +19,23 @@ #include #include #include -#include "utils/check_convert_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAttention = "Attention"; /// \brief MultiHead-Attention op in MindIR. -class MS_CORE_API Attention : public PrimitiveC { +class MIND_API Attention : public BaseOperator { public: + MIND_API_BASE_MEMBER(Attention); /// \brief Constructor. - Attention() : PrimitiveC(kNameAttention) { + Attention() : BaseOperator(kNameAttention) { InitIOName( {"q", "k", "v", "weight_q", "weight_k", "weight_v", "weight_o", "bias_q", "bias_k", "bias_v", "bias_o", "mask"}, {"output"}); } - /// \brief Destructor. - ~Attention() override = default; - MS_DECLARE_PARENT(Attention, PrimitiveC); /// \brief Initialize Attention op. void Init() const {} }; diff --git a/mindspore/core/ops/audio_spectrogram.cc b/mindspore/core/ops/audio_spectrogram.cc index 020298a6c6..bb7a3a65d1 100644 --- a/mindspore/core/ops/audio_spectrogram.cc +++ b/mindspore/core/ops/audio_spectrogram.cc @@ -23,24 +23,28 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(AudioSpectrogram, PrimitiveC, BaseOperator); void AudioSpectrogram::set_window_size(const int64_t window_size) { - (void)this->AddAttr(kWindowSize, MakeValue(window_size)); + (void)this->AddAttr(kWindowSize, api::MakeValue(window_size)); } int64_t AudioSpectrogram::get_window_size() const { auto value_ptr = GetAttr(kWindowSize); return GetValue(value_ptr); } -void AudioSpectrogram::set_stride(const int64_t stride) { (void)this->AddAttr(kStride, MakeValue(stride)); } +void AudioSpectrogram::set_stride(const int64_t stride) { (void)this->AddAttr(kStride, api::MakeValue(stride)); } int64_t AudioSpectrogram::get_stride() const { auto value_ptr = GetAttr(kStride); return GetValue(value_ptr); } -void AudioSpectrogram::set_mag_square(const bool mag_square) { (void)this->AddAttr(kMagSquare, MakeValue(mag_square)); } +void AudioSpectrogram::set_mag_square(const bool mag_square) { + (void)this->AddAttr(kMagSquare, api::MakeValue(mag_square)); +} bool AudioSpectrogram::get_mag_square() const { auto value_ptr = GetAttr(kMagSquare); return GetValue(value_ptr); diff --git a/mindspore/core/ops/audio_spectrogram.h b/mindspore/core/ops/audio_spectrogram.h index e67c4fac6b..14a20dcc85 100644 --- a/mindspore/core/ops/audio_spectrogram.h +++ b/mindspore/core/ops/audio_spectrogram.h @@ -20,23 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAudioSpectrogram = "AudioSpectrogram"; /// \brief AudioSpectrogram defined AudioSpectrogram operator prototype. -class MS_CORE_API AudioSpectrogram : public PrimitiveC { +class MIND_API AudioSpectrogram : public BaseOperator { public: + MIND_API_BASE_MEMBER(AudioSpectrogram); /// \brief Constructor. - AudioSpectrogram() : PrimitiveC(kNameAudioSpectrogram) {} - - /// \brief Destructor. - ~AudioSpectrogram() = default; - - MS_DECLARE_PARENT(AudioSpectrogram, PrimitiveC); + AudioSpectrogram() : BaseOperator(kNameAudioSpectrogram) {} /// \brief Method to init the op's attributes. /// @@ -75,8 +70,8 @@ class MS_CORE_API AudioSpectrogram : public PrimitiveC { /// \return a boolean value. bool get_mag_square() const; }; -AbstractBasePtr AudioSpectrogramInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AudioSpectrogramInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/avg_pool.cc b/mindspore/core/ops/avg_pool.cc index 5b5a509da4..cd5dceb0f6 100644 --- a/mindspore/core/ops/avg_pool.cc +++ b/mindspore/core/ops/avg_pool.cc @@ -23,35 +23,37 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void AvgPool::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode AvgPool::get_pad_mode() const { return PadMode(GetValue(GetAttr(kPadMode))); } void AvgPool::set_kernel_size(const std::vector &kernel_size) { - (void)this->AddAttr(kKernelSize, - MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, this->name()))); + (void)this->AddAttr( + kKernelSize, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, this->name()))); } std::vector AvgPool::get_kernel_size() const { return GetValue>(GetAttr(kKernelSize)); } void AvgPool::set_strides(const std::vector &strides) { - (void)this->AddAttr(kStrides, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStrides, strides, this->name()))); + (void)this->AddAttr(kStrides, + api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStrides, strides, this->name()))); } std::vector AvgPool::get_strides() const { return GetValue>(GetAttr(kStrides)); } void AvgPool::set_format(const Format &format) { int64_t f = format; - (void)this->AddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } Format AvgPool::get_format() const { return Format(GetValue(GetAttr(kFormat))); } -void AvgPool::set_pad(const std::vector &pad) { (void)this->AddAttr(kPad, MakeValue(pad)); } +void AvgPool::set_pad(const std::vector &pad) { (void)this->AddAttr(kPad, api::MakeValue(pad)); } std::vector AvgPool::get_pad() const { auto value_ptr = GetAttr(kPad); @@ -60,7 +62,7 @@ std::vector AvgPool::get_pad() const { void AvgPool::set_round_mode(const RoundMode &round_mode) { int64_t swi = round_mode; - (void)this->AddAttr(kRoundMode, MakeValue(swi)); + (void)this->AddAttr(kRoundMode, api::MakeValue(swi)); } RoundMode AvgPool::get_round_mode() const { @@ -78,6 +80,7 @@ void AvgPool::Init(const std::vector &kernel_size, const std::vectorset_round_mode(round_mode); } +MIND_API_BASE_IMPL(AvgPool, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameAvgPool, AvgPool); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/avg_pool.h b/mindspore/core/ops/avg_pool.h index 888139baa2..c66084c6af 100644 --- a/mindspore/core/ops/avg_pool.h +++ b/mindspore/core/ops/avg_pool.h @@ -21,22 +21,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameAvgPool = "AvgPool"; /// \brief Average pooling operation. Refer to Python API @ref mindspore.ops.AvgPool for more details. -class MS_CORE_API AvgPool : public PrimitiveC { +class MIND_API AvgPool : public BaseOperator { public: + MIND_API_BASE_MEMBER(AvgPool); /// \brief Constructor. - AvgPool() : PrimitiveC(kNameAvgPool) { InitIOName({"x"}, {"output"}); } - explicit AvgPool(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~AvgPool() = default; - MS_DECLARE_PARENT(AvgPool, PrimitiveC); + AvgPool() : BaseOperator(kNameAvgPool) { InitIOName({"x"}, {"output"}); } + explicit AvgPool(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.AvgPool for the inputs. void Init(const std::vector &kernel_size = {1}, const std::vector &stride = {1}, const PadMode &pad_mode = VALID, const Format &format = NCHW, @@ -80,8 +78,8 @@ class MS_CORE_API AvgPool : public PrimitiveC { RoundMode get_round_mode() const; }; -AbstractBasePtr AvgPoolInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AvgPoolInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/avg_pool_3d.cc b/mindspore/core/ops/avg_pool_3d.cc index 9a87502f3f..4a22a58768 100644 --- a/mindspore/core/ops/avg_pool_3d.cc +++ b/mindspore/core/ops/avg_pool_3d.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -180,6 +181,7 @@ TypePtr AvgPool3DInferType(const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(AvgPool3DInferShape(primitive, input_args), AvgPool3DInferType(primitive, input_args)); diff --git a/mindspore/core/ops/avg_pool_3d.h b/mindspore/core/ops/avg_pool_3d.h index f0e86b0703..87b7a06c42 100644 --- a/mindspore/core/ops/avg_pool_3d.h +++ b/mindspore/core/ops/avg_pool_3d.h @@ -21,24 +21,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief 3D Average pooling operation. Refer to Python API @ref mindspore.ops.AvgPool3D for more details. -class MS_CORE_API AvgPool3D : public PrimitiveC { +class MIND_API AvgPool3D : public BaseOperator { public: + MIND_API_BASE_MEMBER(AvgPool3D); /// \brief Constructor. - AvgPool3D() : PrimitiveC(prim::kPrimAvgPool3D->name()) { InitIOName({"input"}, {"output"}); } - /// \brief Destructor. - ~AvgPool3D() = default; - MS_DECLARE_PARENT(AvgPool3D, PrimitiveC); + AvgPool3D() : BaseOperator("AvgPool3D") { InitIOName({"input"}, {"output"}); } }; -AbstractBasePtr AvgPool3DInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AvgPool3DInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/base_operator.cc b/mindspore/core/ops/base_operator.cc index a54ee4e56d..bc6f42c2fc 100644 --- a/mindspore/core/ops/base_operator.cc +++ b/mindspore/core/ops/base_operator.cc @@ -15,13 +15,18 @@ */ #include "ops/base_operator.h" - #include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(BaseOperator, PrimitiveC, api::Primitive); BaseOperator::BaseOperator(const std::string &name) : api::Primitive(std::make_shared(name)) {} +PrimitiveCPtr BaseOperator::GetPrim() { + PrimitiveCPtr res = std::dynamic_pointer_cast(impl_); + return res; +} void BaseOperator::InitIOName(const std::vector &inputs_name, const std::vector &outputs_name) { (void)AddAttr("input_names", api::MakeValue(inputs_name)); diff --git a/mindspore/core/ops/base_operator.h b/mindspore/core/ops/base_operator.h index 821d6172fc..070c999efe 100644 --- a/mindspore/core/ops/base_operator.h +++ b/mindspore/core/ops/base_operator.h @@ -17,26 +17,36 @@ #ifndef MINDSPORE_CORE_OPS_BASE_OPERATOR_ #define MINDSPORE_CORE_OPS_BASE_OPERATOR_ -#include #include +#include #include #include "mindapi/ir/primitive.h" +namespace mindspore { namespace abstract { class AnalysisEngine; using AnalysisEnginePtr = std::shared_ptr; class AbstractBase; -using AbstractBasePtr = std::shared_ptr; +using AbstractBasePtr = std::shared_ptr; } // namespace abstract +} // namespace mindspore + +namespace mindspore { +class Primitive; +using PrimitivePtr = std::shared_ptr; +} // namespace mindspore namespace mindspore { namespace ops { -class BaseOperator : public api::Primitive { +class PrimitiveC; +using PrimitiveCPtr = std::shared_ptr; +class MIND_API BaseOperator : public api::Primitive { public: + MIND_API_BASE_MEMBER(BaseOperator); explicit BaseOperator(const std::string &name); - ~BaseOperator() = default; + PrimitiveCPtr GetPrim(); protected: void InitIOName(const std::vector &inputs_name, const std::vector &outputs_name); diff --git a/mindspore/core/ops/batch_matmul.cc b/mindspore/core/ops/batch_matmul.cc index 474bd16922..82473b8939 100644 --- a/mindspore/core/ops/batch_matmul.cc +++ b/mindspore/core/ops/batch_matmul.cc @@ -23,6 +23,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -134,14 +135,15 @@ TypePtr BatchMatmulInferType(const PrimitivePtr &prim, const std::vector #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Computes matrix multiplication between two tensors by batch. /// Refer to Python API @ref mindspore.ops.BatchMatmul for more details. -class MS_CORE_API BatchMatmul : public PrimitiveC { +class MIND_API BatchMatmul : public BaseOperator { public: + MIND_API_BASE_MEMBER(BatchMatmul); /// \brief Constructor. - BatchMatmul() : PrimitiveC(prim::kPrimBatchMatMul->name()) { InitIOName({"x1", "x2"}, {"output"}); } - /// \brief Destructor. - ~BatchMatmul() = default; - MS_DECLARE_PARENT(BatchMatmul, PrimitiveC); + BatchMatmul() : BaseOperator("BatchMatMul") { InitIOName({"x1", "x2"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BatchMatmul for the inputs. void Init(bool transpose_a = false, bool transpose_b = false); /// \brief Set transpose_a. @@ -48,8 +45,8 @@ class MS_CORE_API BatchMatmul : public PrimitiveC { /// \return transpose_b. bool get_transpose_b() const; }; -AbstractBasePtr BatchMatmulInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BatchMatmulInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/batch_norm.cc b/mindspore/core/ops/batch_norm.cc index 2d48d8129e..cfdafaca07 100644 --- a/mindspore/core/ops/batch_norm.cc +++ b/mindspore/core/ops/batch_norm.cc @@ -21,9 +21,12 @@ #include "ops/batch_norm.h" #include "abstract/primitive_infer_map.h" #include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(BatchNorm, PrimitiveC, BaseOperator); void BatchNorm::Init(const bool is_training, const float epsilon, const float momentum, const Format &format) { set_is_training(is_training); set_epsilon(epsilon); @@ -31,21 +34,23 @@ void BatchNorm::Init(const bool is_training, const float epsilon, const float mo set_momentum(momentum); } -void BatchNorm::set_is_training(const bool is_training) { (void)this->AddAttr(kIsTraining, MakeValue(is_training)); } +void BatchNorm::set_is_training(const bool is_training) { + (void)this->AddAttr(kIsTraining, api::MakeValue(is_training)); +} void BatchNorm::set_epsilon(const float epsilon) { CheckAndConvertUtils::CheckInRange(kEpsilon, epsilon, kIncludeBoth, {0.0, 1.0}, this->name()); - (void)this->AddAttr(kEpsilon, MakeValue(epsilon)); + (void)this->AddAttr(kEpsilon, api::MakeValue(epsilon)); } void BatchNorm::set_format(const Format &format) { int64_t f = format; - (void)this->AddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } void BatchNorm::set_momentum(const float momentun) { CheckAndConvertUtils::CheckInRange(kMomentum, momentun, kIncludeBoth, {0.0, 1.0}, this->name()); - (void)this->AddAttr(kMomentum, MakeValue(momentun)); + (void)this->AddAttr(kMomentum, api::MakeValue(momentun)); } float BatchNorm::get_momentum() const { diff --git a/mindspore/core/ops/batch_norm.h b/mindspore/core/ops/batch_norm.h index f28dbb5551..50dc18a940 100644 --- a/mindspore/core/ops/batch_norm.h +++ b/mindspore/core/ops/batch_norm.h @@ -20,25 +20,22 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameBatchNorm = "BatchNorm"; /// \brief Batch Normalization for input data and updated parameters. /// Refer to Python API @ref mindspore.ops.BatchNorm for more details. -class MS_CORE_API BatchNorm : public PrimitiveC { +class MIND_API BatchNorm : public BaseOperator { public: + MIND_API_BASE_MEMBER(BatchNorm); /// \brief Constructor. - BatchNorm() : PrimitiveC(kNameBatchNorm) { + BatchNorm() : BaseOperator(kNameBatchNorm) { InitIOName({"x", "scale", "offset", "mean", "variance"}, {"y", "batch_mean", "batch_variance", "reserve_space_1", "reserve_space_2"}); } - /// \brief Destructor. - ~BatchNorm() = default; - MS_DECLARE_PARENT(BatchNorm, PrimitiveC); /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BatchNorm for the inputs. void Init(const bool is_training = false, const float epsilon = 1e-5, const float momentun = 0.1, const Format &format = NCHW); @@ -68,8 +65,8 @@ class MS_CORE_API BatchNorm : public PrimitiveC { float get_momentum() const; }; -AbstractBasePtr BatchNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BatchNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimBatchNormPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/batch_to_space.cc b/mindspore/core/ops/batch_to_space.cc index 86899bac3a..3881838207 100644 --- a/mindspore/core/ops/batch_to_space.cc +++ b/mindspore/core/ops/batch_to_space.cc @@ -18,16 +18,18 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(BatchToSpace, PrimitiveC, BaseOperator); void BatchToSpace::Init(const std::vector &block_size, const std::vector> &crops) { this->set_block_size(block_size); this->set_crops(crops); } void BatchToSpace::set_block_size(const std::vector &block_size) { - (void)this->AddAttr(kBlockSize, MakeValue(block_size)); + (void)this->AddAttr(kBlockSize, api::MakeValue(block_size)); } std::vector BatchToSpace::get_block_size() const { @@ -36,7 +38,7 @@ std::vector BatchToSpace::get_block_size() const { } void BatchToSpace::set_crops(const std::vector> &crops) { - (void)this->AddAttr(kCrops, MakeValue(crops)); + (void)this->AddAttr(kCrops, api::MakeValue(crops)); } std::vector> BatchToSpace::get_crops() const { diff --git a/mindspore/core/ops/batch_to_space.h b/mindspore/core/ops/batch_to_space.h index d78eefab3a..0342fd70bd 100644 --- a/mindspore/core/ops/batch_to_space.h +++ b/mindspore/core/ops/batch_to_space.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBatchToSpace = "BatchToSpace"; /// \brief Divides batch dimension with blocks and interleaves these blocks back into spatial dimensions. /// Refer to Python API @ref mindspore.ops.BatchToSpace for more details. -class MS_CORE_API BatchToSpace : public PrimitiveC { +class MIND_API BatchToSpace : public BaseOperator { public: + MIND_API_BASE_MEMBER(BatchToSpace); /// \brief Constructor. - BatchToSpace() : PrimitiveC(kNameBatchToSpace) {} - /// \brief Destructor. - ~BatchToSpace() = default; - MS_DECLARE_PARENT(BatchToSpace, PrimitiveC); + BatchToSpace() : BaseOperator(kNameBatchToSpace) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BatchToSpace for the inputs. void Init(const std::vector &block_size, const std::vector> &crops); /// \brief Set block_size. @@ -51,8 +48,8 @@ class MS_CORE_API BatchToSpace : public PrimitiveC { std::vector> get_crops() const; }; -AbstractBasePtr BatchToSpaceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BatchToSpaceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/batch_to_space_nd.cc b/mindspore/core/ops/batch_to_space_nd.cc index cbc63945df..48766dd4f5 100644 --- a/mindspore/core/ops/batch_to_space_nd.cc +++ b/mindspore/core/ops/batch_to_space_nd.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -94,6 +95,7 @@ TypePtr BatchToSpaceNDInferType(const std::vector &input_args) } } // namespace +MIND_API_BASE_IMPL(BatchToSpaceND, PrimitiveC, BaseOperator); void BatchToSpaceND::set_crops(std::vector> crops) { const int64_t crop_size = 2; (void)CheckAndConvertUtils::CheckInteger(kCrops, SizeToLong(crops.size()), kEqual, crop_size, this->name()); @@ -106,7 +108,7 @@ void BatchToSpaceND::set_crops(std::vector> crops) { (void)CheckAndConvertUtils::CheckInteger(kCrops, crops[i][j], kGreaterEqual, 0, this->name()); } } - (void)this->AddAttr(kCrops, MakeValue(crops)); + (void)this->AddAttr(kCrops, api::MakeValue(crops)); } std::vector> BatchToSpaceND::get_crops() const { @@ -120,7 +122,7 @@ void BatchToSpaceND::set_block_shape(std::vector block_shape) { for (size_t i = 0; i < block_shape.size(); i++) { (void)CheckAndConvertUtils::CheckInteger(kBlockShape, block_shape[i], kGreaterEqual, 1, this->name()); } - (void)this->AddAttr(kBlockShape, MakeValue(block_shape)); + (void)this->AddAttr(kBlockShape, api::MakeValue(block_shape)); } std::vector BatchToSpaceND::get_block_shape() const { diff --git a/mindspore/core/ops/batch_to_space_nd.h b/mindspore/core/ops/batch_to_space_nd.h index c4dfdd9454..2dcfc2c81b 100644 --- a/mindspore/core/ops/batch_to_space_nd.h +++ b/mindspore/core/ops/batch_to_space_nd.h @@ -21,22 +21,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBatchToSpaceND = "BatchToSpaceND"; /// \brief Divides batch dimension with blocks and interleaves these blocks back into spatial dimensions. /// Refer to Python API @ref mindspore.ops.BatchToSpaceND for more details. -class MS_CORE_API BatchToSpaceND : public PrimitiveC { +class MIND_API BatchToSpaceND : public BaseOperator { public: + MIND_API_BASE_MEMBER(BatchToSpaceND); /// \brief Constructor. - BatchToSpaceND() : PrimitiveC(kNameBatchToSpaceND) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~BatchToSpaceND() = default; - MS_DECLARE_PARENT(BatchToSpaceND, PrimitiveC); + BatchToSpaceND() : BaseOperator(kNameBatchToSpaceND) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BatchToSpaceND for the inputs. void Init(const std::vector block_shape, const std::vector> crops); /// \brief Set crops. @@ -52,8 +49,8 @@ class MS_CORE_API BatchToSpaceND : public PrimitiveC { /// \return crops. std::vector> get_crops() const; }; -AbstractBasePtr BatchToSpaceNDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BatchToSpaceNDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBatchToSpaceNDPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/bessel_i0.cc b/mindspore/core/ops/bessel_i0.cc index 03d77b6f00..c21f7163e7 100644 --- a/mindspore/core/ops/bessel_i0.cc +++ b/mindspore/core/ops/bessel_i0.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -37,6 +38,7 @@ TypePtr BesselI0InferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bessel_i0.h b/mindspore/core/ops/bessel_i0.h index 31baec02bd..2605080d8b 100644 --- a/mindspore/core/ops/bessel_i0.h +++ b/mindspore/core/ops/bessel_i0.h @@ -21,22 +21,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBesselI0 = "BesselI0"; -class BesselI0 : public PrimitiveC { +class MIND_API BesselI0 : public BaseOperator { public: - BesselI0() : PrimitiveC(kNameBesselI0) { InitIOName({"x"}, {"y"}); } - ~BesselI0() = default; - MS_DECLARE_PARENT(BesselI0, PrimitiveC); + MIND_API_BASE_MEMBER(BesselI0); + BesselI0() : BaseOperator(kNameBesselI0) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr BesselI0Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BesselI0Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_Bessel_I0_H_ diff --git a/mindspore/core/ops/bessel_i0e.cc b/mindspore/core/ops/bessel_i0e.cc index 0459ac8014..f618b5f600 100644 --- a/mindspore/core/ops/bessel_i0e.cc +++ b/mindspore/core/ops/bessel_i0e.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,8 @@ TypePtr BesselI0eInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bessel_i0e.h b/mindspore/core/ops/bessel_i0e.h index 06312cbe7c..6e9157cf69 100644 --- a/mindspore/core/ops/bessel_i0e.h +++ b/mindspore/core/ops/bessel_i0e.h @@ -19,19 +19,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBesselI0e = "BesselI0e"; -class BesselI0e : public PrimitiveC { +class MIND_API BesselI0e : public BaseOperator { public: - BesselI0e() : PrimitiveC(kNameBesselI0e) { InitIOName({"x"}, {"output"}); } - ~BesselI0e() = default; - MS_DECLARE_PARENT(BesselI0e, PrimitiveC); + MIND_API_BASE_MEMBER(BesselI0e); + BesselI0e() : BaseOperator(kNameBesselI0e) { InitIOName({"x"}, {"output"}); } void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/bessel_i1.cc b/mindspore/core/ops/bessel_i1.cc index 0326cf7c9c..9e1f20cfcd 100644 --- a/mindspore/core/ops/bessel_i1.cc +++ b/mindspore/core/ops/bessel_i1.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -37,6 +38,7 @@ TypePtr BesselI1InferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bessel_i1.h b/mindspore/core/ops/bessel_i1.h index 4b1a391763..d61ebb2d5f 100644 --- a/mindspore/core/ops/bessel_i1.h +++ b/mindspore/core/ops/bessel_i1.h @@ -21,22 +21,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBesselI1 = "BesselI1"; -class BesselI1 : public PrimitiveC { +class MIND_API BesselI1 : public BaseOperator { public: - BesselI1() : PrimitiveC(kNameBesselI1) { InitIOName({"x"}, {"y"}); } - ~BesselI1() = default; - MS_DECLARE_PARENT(BesselI1, PrimitiveC); + MIND_API_BASE_MEMBER(BesselI1); + BesselI1() : BaseOperator(kNameBesselI1) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr BesselI1Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BesselI1Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_Bessel_I1_H_ diff --git a/mindspore/core/ops/bessel_i1e.cc b/mindspore/core/ops/bessel_i1e.cc index 88cc91e87a..0fa548f069 100644 --- a/mindspore/core/ops/bessel_i1e.cc +++ b/mindspore/core/ops/bessel_i1e.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,8 @@ TypePtr BesselI1eInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bessel_i1e.h b/mindspore/core/ops/bessel_i1e.h index b8266ab1db..f60154bdff 100644 --- a/mindspore/core/ops/bessel_i1e.h +++ b/mindspore/core/ops/bessel_i1e.h @@ -19,19 +19,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBesselI1e = "BesselI1e"; -class BesselI1e : public PrimitiveC { +class MIND_API BesselI1e : public BaseOperator { public: - BesselI1e() : PrimitiveC(kNameBesselI1e) { InitIOName({"x"}, {"output"}); } - ~BesselI1e() = default; - MS_DECLARE_PARENT(BesselI1e, PrimitiveC); + MIND_API_BASE_MEMBER(BesselI1e); + BesselI1e() : BaseOperator(kNameBesselI1e) { InitIOName({"x"}, {"output"}); } void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/bias_add.cc b/mindspore/core/ops/bias_add.cc index e017a38769..965d7cf7db 100644 --- a/mindspore/core/ops/bias_add.cc +++ b/mindspore/core/ops/bias_add.cc @@ -24,6 +24,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/ms_context.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -97,9 +98,11 @@ TypePtr BiasAddInferType(const PrimitivePtr &prim, const std::vectorAddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } Format BiasAdd::get_format() const { auto value_ptr = GetAttr(kFormat); diff --git a/mindspore/core/ops/bias_add.h b/mindspore/core/ops/bias_add.h index dbfdd978ad..d7481bf208 100644 --- a/mindspore/core/ops/bias_add.h +++ b/mindspore/core/ops/bias_add.h @@ -20,23 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -// Add -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { -constexpr auto kNameBiasAdd = prim::kBiasAdd; +constexpr auto kNameBiasAdd = "BiasAdd"; /// \brief Returns sum of input and bias tensor. Refer to Python API @ref mindspore.ops.BiasAdd for more details. -class MS_CORE_API BiasAdd : public PrimitiveC { +class MIND_API BiasAdd : public BaseOperator { public: + MIND_API_BASE_MEMBER(BiasAdd); /// \brief Constructor. - BiasAdd() : PrimitiveC(prim::kPrimBiasAdd->name()) { InitIOName({"x", "b"}, {"output"}); } - /// \brief Destructor. - ~BiasAdd() = default; - MS_DECLARE_PARENT(BiasAdd, PrimitiveC); + BiasAdd() : BaseOperator("BiasAdd") { InitIOName({"x", "b"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BiasAdd for the inputs. void Init(const Format &format = NCHW); /// \brief Set format. @@ -46,8 +42,8 @@ class MS_CORE_API BiasAdd : public PrimitiveC { /// \return format. Format get_format() const; }; -AbstractBasePtr BiasAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BiasAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/binary_cross_entropy.cc b/mindspore/core/ops/binary_cross_entropy.cc index ddc0bd93dc..57855b51e3 100644 --- a/mindspore/core/ops/binary_cross_entropy.cc +++ b/mindspore/core/ops/binary_cross_entropy.cc @@ -25,6 +25,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -80,9 +81,10 @@ TypePtr BinaryCrossEntroyInferType(const PrimitivePtr &prim, const std::vectorAddAttr(kReduction, MakeValue(swi)); + (void)this->AddAttr(kReduction, api::MakeValue(swi)); } Reduction BinaryCrossEntropy::get_reduction() const { diff --git a/mindspore/core/ops/binary_cross_entropy.h b/mindspore/core/ops/binary_cross_entropy.h index 048ebb45bd..6f9dc56355 100644 --- a/mindspore/core/ops/binary_cross_entropy.h +++ b/mindspore/core/ops/binary_cross_entropy.h @@ -20,23 +20,19 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBinaryCrossEntropy = "BinaryCrossEntropy"; /// \brief Computes the binary cross entropy between the logits and the labels. /// Refer to Python API @ref mindspore.ops.BinaryCrossEntropy for more details. -class MS_CORE_API BinaryCrossEntropy : public PrimitiveC { +class MIND_API BinaryCrossEntropy : public BaseOperator { public: + MIND_API_BASE_MEMBER(BinaryCrossEntropy); /// \brief Constructor. - BinaryCrossEntropy() : PrimitiveC(kNameBinaryCrossEntropy) { InitIOName({"x", "y", "weight"}, {"output"}); } - /// \brief Destructor. - ~BinaryCrossEntropy() = default; - MS_DECLARE_PARENT(BinaryCrossEntropy, PrimitiveC); + BinaryCrossEntropy() : BaseOperator(kNameBinaryCrossEntropy) { InitIOName({"x", "y", "weight"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BinaryCrossEntropy for the inputs. void Init(const Reduction &reduction = MEAN); /// \brief Set reduction. @@ -46,8 +42,9 @@ class MS_CORE_API BinaryCrossEntropy : public PrimitiveC { /// \return reduction. Reduction get_reduction() const; }; -AbstractBasePtr BinaryCrossEntropyGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BinaryCrossEntropyGradInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBinaryCrossEntropyPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/bitwiseand.cc b/mindspore/core/ops/bitwiseand.cc index 4479f0b0bc..52380b8284 100644 --- a/mindspore/core/ops/bitwiseand.cc +++ b/mindspore/core/ops/bitwiseand.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -47,6 +48,8 @@ TypePtr BitwiseAndInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(BitwiseAnd, PrimitiveC, BaseOperator); AbstractBasePtr BitwiseAndInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto op_name = primitive->name(); diff --git a/mindspore/core/ops/bitwiseand.h b/mindspore/core/ops/bitwiseand.h index 9071e1d942..7a80ba2511 100644 --- a/mindspore/core/ops/bitwiseand.h +++ b/mindspore/core/ops/bitwiseand.h @@ -20,23 +20,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBitwiseAnd = "BitwiseAnd"; -class BitwiseAnd : public PrimitiveC { +class MIND_API BitwiseAnd : public BaseOperator { public: - BitwiseAnd() : PrimitiveC(kNameBitwiseAnd) { InitIOName({"x1", "x2"}, {"y"}); } - explicit BitwiseAnd(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x1", "x2"}, {"y"}); } - ~BitwiseAnd() = default; - MS_DECLARE_PARENT(BitwiseAnd, PrimitiveC); + MIND_API_BASE_MEMBER(BitwiseAnd); + BitwiseAnd() : BaseOperator(kNameBitwiseAnd) { InitIOName({"x1", "x2"}, {"y"}); } + explicit BitwiseAnd(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x1", "x2"}, {"y"}); } void Init() {} }; -AbstractBasePtr BitwiseAndInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BitwiseAndInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBitwiseAndPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/bitwiseor.cc b/mindspore/core/ops/bitwiseor.cc index f27033d44a..e7b5d12e8d 100644 --- a/mindspore/core/ops/bitwiseor.cc +++ b/mindspore/core/ops/bitwiseor.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -47,6 +48,8 @@ TypePtr BitwiseOrInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(BitwiseOr, PrimitiveC, BaseOperator); AbstractBasePtr BitwiseOrInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto op_name = primitive->name(); diff --git a/mindspore/core/ops/bitwiseor.h b/mindspore/core/ops/bitwiseor.h index 589bd7dc9c..8dadc423da 100644 --- a/mindspore/core/ops/bitwiseor.h +++ b/mindspore/core/ops/bitwiseor.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBitwiseOr = "BitwiseOr"; -class BitwiseOr : public PrimitiveC { +class MIND_API BitwiseOr : public BaseOperator { public: - BitwiseOr() : PrimitiveC(kNameBitwiseOr) { InitIOName({"x1", "x2"}, {"y"}); } - explicit BitwiseOr(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x1", "x2"}, {"y"}); } - ~BitwiseOr() = default; - MS_DECLARE_PARENT(BitwiseOr, PrimitiveC); + MIND_API_BASE_MEMBER(BitwiseOr); + BitwiseOr() : BaseOperator(kNameBitwiseOr) { InitIOName({"x1", "x2"}, {"y"}); } + explicit BitwiseOr(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x1", "x2"}, {"y"}); } }; -AbstractBasePtr BitwiseOrInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BitwiseOrInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBitwiseOrPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/bitwisexor.cc b/mindspore/core/ops/bitwisexor.cc index 8886a26ced..d6aa548be5 100644 --- a/mindspore/core/ops/bitwisexor.cc +++ b/mindspore/core/ops/bitwisexor.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -47,6 +48,8 @@ TypePtr BitwiseXorInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(BitwiseXor, PrimitiveC, BaseOperator); AbstractBasePtr BitwiseXorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto op_name = primitive->name(); diff --git a/mindspore/core/ops/bitwisexor.h b/mindspore/core/ops/bitwisexor.h index e276efde67..2877b2a859 100644 --- a/mindspore/core/ops/bitwisexor.h +++ b/mindspore/core/ops/bitwisexor.h @@ -20,23 +20,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBitwiseXor = "BitwiseXor"; -class BitwiseXor : public PrimitiveC { +class MIND_API BitwiseXor : public BaseOperator { public: - BitwiseXor() : PrimitiveC(kNameBitwiseXor) { InitIOName({"x1", "x2"}, {"y"}); } - explicit BitwiseXor(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x1", "x2"}, {"y"}); } - ~BitwiseXor() = default; - MS_DECLARE_PARENT(BitwiseXor, PrimitiveC); + MIND_API_BASE_MEMBER(BitwiseXor); + BitwiseXor() : BaseOperator(kNameBitwiseXor) { InitIOName({"x1", "x2"}, {"y"}); } + explicit BitwiseXor(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x1", "x2"}, {"y"}); } void Init() {} }; -AbstractBasePtr BitwiseXorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BitwiseXorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBitwiseXorPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/bn_training_reduce.cc b/mindspore/core/ops/bn_training_reduce.cc index d099731cee..d024fbe888 100644 --- a/mindspore/core/ops/bn_training_reduce.cc +++ b/mindspore/core/ops/bn_training_reduce.cc @@ -24,6 +24,7 @@ #include "utils/tensor_construct_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -77,6 +78,8 @@ TypePtr BNTrainingReduceInferType(const PrimitivePtr &primitive, const std::vect return std::make_shared(std::vector{input_type, input_type}); } } // namespace + +MIND_API_BASE_IMPL(BNTrainingReduce, PrimitiveC, BaseOperator); AbstractBasePtr BNTrainingReduceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bn_training_reduce.h b/mindspore/core/ops/bn_training_reduce.h index 9d8dfe9d88..e5181202f9 100644 --- a/mindspore/core/ops/bn_training_reduce.h +++ b/mindspore/core/ops/bn_training_reduce.h @@ -19,21 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamekBNTrainingReduce = "BNTrainingReduce"; -class BNTrainingReduce : public PrimitiveC { +class MIND_API BNTrainingReduce : public BaseOperator { public: - BNTrainingReduce() : PrimitiveC(kNamekBNTrainingReduce) { InitIOName({"x"}, {"sum", "square_sum"}); } - ~BNTrainingReduce() = default; - MS_DECLARE_PARENT(BNTrainingReduce, PrimitiveC); + MIND_API_BASE_MEMBER(BNTrainingReduce); + BNTrainingReduce() : BaseOperator(kNamekBNTrainingReduce) { InitIOName({"x"}, {"sum", "square_sum"}); } }; -AbstractBasePtr BNTrainingReduceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BNTrainingReduceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBNTrainingReduce = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/bn_training_update.cc b/mindspore/core/ops/bn_training_update.cc index a34fde4d3b..3ae592e1fc 100644 --- a/mindspore/core/ops/bn_training_update.cc +++ b/mindspore/core/ops/bn_training_update.cc @@ -22,6 +22,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -122,6 +124,7 @@ TuplePtr BNTrainingUpdateInferType(const PrimitivePtr &primitive, const std::vec } } // namespace +MIND_API_BASE_IMPL(BNTrainingUpdate, PrimitiveC, BaseOperator); AbstractBasePtr BNTrainingUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bn_training_update.h b/mindspore/core/ops/bn_training_update.h index 50725455bf..a63664fc6f 100644 --- a/mindspore/core/ops/bn_training_update.h +++ b/mindspore/core/ops/bn_training_update.h @@ -22,25 +22,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBNTrainingUpdate = "BNTrainingUpdate"; -class MS_CORE_API BNTrainingUpdate : public PrimitiveC { +class MIND_API BNTrainingUpdate : public BaseOperator { public: - BNTrainingUpdate() : PrimitiveC(kNameBNTrainingUpdate) { + MIND_API_BASE_MEMBER(BNTrainingUpdate); + BNTrainingUpdate() : BaseOperator(kNameBNTrainingUpdate) { InitIOName({"x", "sum", "square_sum", "scale", "b", "mean", "variance"}, {"y", "running_mean", "running_variance", "save_mean", "save_inv_variance"}); } - ~BNTrainingUpdate() = default; - MS_DECLARE_PARENT(BNTrainingUpdate, PrimitiveC); }; -AbstractBasePtr BNTrainingUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BNTrainingUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBNTrainingUpdatePtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/bounding_box_decode.cc b/mindspore/core/ops/bounding_box_decode.cc index 87e3a2ee36..0f95690be3 100644 --- a/mindspore/core/ops/bounding_box_decode.cc +++ b/mindspore/core/ops/bounding_box_decode.cc @@ -21,6 +21,8 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -94,6 +96,7 @@ TypePtr BoundingBoxDecodeInferType(const PrimitivePtr &primitive, const std::vec } } // namespace +MIND_API_BASE_IMPL(BoundingBoxDecode, PrimitiveC, BaseOperator); AbstractBasePtr BoundingBoxDecodeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/bounding_box_decode.h b/mindspore/core/ops/bounding_box_decode.h index 7f668c682a..885b267a1f 100644 --- a/mindspore/core/ops/bounding_box_decode.h +++ b/mindspore/core/ops/bounding_box_decode.h @@ -22,23 +22,20 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBoundingBoxDecode = "BoundingBoxDecode"; -class MS_CORE_API BoundingBoxDecode : public PrimitiveC { +class MIND_API BoundingBoxDecode : public BaseOperator { public: - BoundingBoxDecode() : PrimitiveC(kNameBoundingBoxDecode) { InitIOName({"anchor_box", "deltas"}, {"output"}); } - ~BoundingBoxDecode() = default; - MS_DECLARE_PARENT(BoundingBoxDecode, PrimitiveC); + MIND_API_BASE_MEMBER(BoundingBoxDecode); + BoundingBoxDecode() : BaseOperator(kNameBoundingBoxDecode) { InitIOName({"anchor_box", "deltas"}, {"output"}); } }; -AbstractBasePtr BoundingBoxDecodeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BoundingBoxDecodeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/broadcast.cc b/mindspore/core/ops/broadcast.cc index 6ba0a6fe0e..1c999be4e4 100644 --- a/mindspore/core/ops/broadcast.cc +++ b/mindspore/core/ops/broadcast.cc @@ -20,18 +20,20 @@ #include "ops/broadcast.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Broadcast, PrimitiveC, BaseOperator); void Broadcast::Init(const int64_t root_rank, const std::string &group) { this->set_root_rank(root_rank); this->set_group(group); } -void Broadcast::set_root_rank(const int64_t root_rank) { (void)this->AddAttr(kKeepProb, MakeValue(root_rank)); } +void Broadcast::set_root_rank(const int64_t root_rank) { (void)this->AddAttr(kKeepProb, api::MakeValue(root_rank)); } void Broadcast::set_group(const std::string &group) { CheckAndConvertUtils::CheckString(kGroup, group, {"hccl_world_group", "hccl_world_group"}, this->name()); - (void)this->AddAttr(kGroup, MakeValue(group)); + (void)this->AddAttr(kGroup, api::MakeValue(group)); } int64_t Broadcast::get_root_rank() const { auto value_ptr = this->GetAttr(kRootRank); diff --git a/mindspore/core/ops/broadcast.h b/mindspore/core/ops/broadcast.h index 3bfdc9a8d2..083b571524 100644 --- a/mindspore/core/ops/broadcast.h +++ b/mindspore/core/ops/broadcast.h @@ -20,21 +20,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBroadcast = "Broadcast"; /// \brief Broadcasts the tensor to the whole group. Refer to Python API @ref mindspore.ops.Broadcast for more details. -class MS_CORE_API Broadcast : public PrimitiveC { +class MIND_API Broadcast : public BaseOperator { public: + MIND_API_BASE_MEMBER(Broadcast); /// \brief Constructor. - Broadcast() : PrimitiveC(kNameBroadcast) {} - /// \brief Destructor. - ~Broadcast() = default; - MS_DECLARE_PARENT(Broadcast, PrimitiveC); + Broadcast() : BaseOperator(kNameBroadcast) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Broadcast for the inputs. void Init(const int64_t root_rank, const std::string &group = "hccl_world_group"); /// \brief Set root_rank. @@ -50,8 +47,8 @@ class MS_CORE_API Broadcast : public PrimitiveC { /// \return group. std::string get_group() const; }; -AbstractBasePtr BroadcastInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BroadcastInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimBroadcast = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/broadcast_to.cc b/mindspore/core/ops/broadcast_to.cc index 21a977ee6d..4d3620fff2 100644 --- a/mindspore/core/ops/broadcast_to.cc +++ b/mindspore/core/ops/broadcast_to.cc @@ -16,6 +16,9 @@ #include #include "ops/broadcast_to.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -73,11 +76,12 @@ TypePtr BroadcastToInferType(const PrimitivePtr &prim, const std::vector &shape) { set_shape(shape); } void BroadcastTo::set_shape(const std::vector &shape) { (void)CheckAndConvertUtils::CheckInteger(kShapeSize, SizeToLong(shape.size()), kGreaterThan, 0, name()); - (void)AddAttr(kShape, MakeValue(shape)); + (void)AddAttr(kShape, api::MakeValue(shape)); } std::vector BroadcastTo::get_shape() const { diff --git a/mindspore/core/ops/broadcast_to.h b/mindspore/core/ops/broadcast_to.h index a4d6e95a35..0007fba8df 100644 --- a/mindspore/core/ops/broadcast_to.h +++ b/mindspore/core/ops/broadcast_to.h @@ -20,22 +20,18 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Broadcasts input tensor to a given shape. /// Refer to Python API @ref mindspore.ops.BroadcastTo for more details. -class MS_CORE_API BroadcastTo : public PrimitiveC { +class MIND_API BroadcastTo : public BaseOperator { public: + MIND_API_BASE_MEMBER(BroadcastTo); /// \brief Constructor. - BroadcastTo() : PrimitiveC(prim::kPrimBroadcastTo->name()) {} - /// \brief Destructor. - ~BroadcastTo() = default; - MS_DECLARE_PARENT(BroadcastTo, PrimitiveC); + BroadcastTo() : BaseOperator("BroadcastTo") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.BroadcastTo for the inputs. void Init(const std::vector &shape); /// \brief Set shape. @@ -46,8 +42,8 @@ class MS_CORE_API BroadcastTo : public PrimitiveC { std::vector get_shape() const; }; -AbstractBasePtr BroadcastToInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BroadcastToInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/call.cc b/mindspore/core/ops/call.cc index 5f6ad18106..d5fcde120c 100644 --- a/mindspore/core/ops/call.cc +++ b/mindspore/core/ops/call.cc @@ -16,9 +16,11 @@ #include "ops/call.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Call, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameCall, Call); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/call.h b/mindspore/core/ops/call.h index f66e462386..3e450ee8ce 100644 --- a/mindspore/core/ops/call.h +++ b/mindspore/core/ops/call.h @@ -16,21 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_CALL_H_ #define MINDSPORE_CORE_OPS_CALL_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCall = "call"; /// \brief Call op means function call in the MindIR. This operator is defined for serialization. -class MS_CORE_API Call : public PrimitiveC { +class MIND_API Call : public BaseOperator { public: + MIND_API_BASE_MEMBER(Call); /// \brief Constructor. - Call() : PrimitiveC(kNameCall) {} - /// \brief Destructor. - ~Call() = default; - MS_DECLARE_PARENT(Call, PrimitiveC); + Call() : BaseOperator(kNameCall) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cast.cc b/mindspore/core/ops/cast.cc index 2babc975be..e431f02dc1 100644 --- a/mindspore/core/ops/cast.cc +++ b/mindspore/core/ops/cast.cc @@ -15,9 +15,12 @@ */ #include "ops/cast.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Cast, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameCast, Cast); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cast.h b/mindspore/core/ops/cast.h index c862544fa4..2694f9dbad 100644 --- a/mindspore/core/ops/cast.h +++ b/mindspore/core/ops/cast.h @@ -19,26 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCast = "Cast"; /// \brief Returns a tensor with the new specified data type. /// Refer to Python API @ref mindspore.ops.Cast for more details. -class MS_CORE_API Cast : public PrimitiveC { +class MIND_API Cast : public BaseOperator { public: + MIND_API_BASE_MEMBER(Cast); /// \brief Constructor. - Cast() : PrimitiveC(kNameCast) { InitIOName({"x", "dst_type"}, {"output"}); } - /// \brief Destructor. - ~Cast() = default; - MS_DECLARE_PARENT(Cast, PrimitiveC); + Cast() : BaseOperator(kNameCast) { InitIOName({"x", "dst_type"}, {"output"}); } }; -AbstractBasePtr CastInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CastInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCast = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cdist.cc b/mindspore/core/ops/cdist.cc index 81ec3c779f..1884560796 100644 --- a/mindspore/core/ops/cdist.cc +++ b/mindspore/core/ops/cdist.cc @@ -19,6 +19,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -57,6 +58,7 @@ TypePtr CdistInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/cdist.h b/mindspore/core/ops/cdist.h index dd97de2b4a..0f3486bf74 100644 --- a/mindspore/core/ops/cdist.h +++ b/mindspore/core/ops/cdist.h @@ -22,26 +22,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCdist = "Cdist"; /// \brief Computes batched the p norm distance between each pair of the two collections of row vectors. /// Refer to Python API @ref mindspore.ops.Cdist for more details. -class MS_CORE_API Cdist : public PrimitiveC { +class MIND_API Cdist : public BaseOperator { public: + MIND_API_BASE_MEMBER(Cdist); /// \brief Constructor. - Cdist() : PrimitiveC(kNameCdist) { InitIOName({"input_x", "input_y"}, {"output"}); } - /// \brief Destructor. - ~Cdist() = default; - MS_DECLARE_PARENT(Cdist, PrimitiveC); + Cdist() : BaseOperator(kNameCdist) { InitIOName({"input_x", "input_y"}, {"output"}); } }; -AbstractBasePtr CdistInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CdistInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/ceil.cc b/mindspore/core/ops/ceil.cc index 52d0e62744..ac74357b7e 100644 --- a/mindspore/core/ops/ceil.cc +++ b/mindspore/core/ops/ceil.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -54,6 +55,7 @@ TypePtr CeilInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = CeilInferType(primitive, input_args); diff --git a/mindspore/core/ops/ceil.h b/mindspore/core/ops/ceil.h index b350099a21..ebcdd0d23c 100644 --- a/mindspore/core/ops/ceil.h +++ b/mindspore/core/ops/ceil.h @@ -19,27 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameCeil = "Ceil"; /// \brief Rounds a tensor up to the closest integer element-wise. /// Refer to Python API @ref mindspore.ops.Ceil for more details. -class MS_CORE_API Ceil : public PrimitiveC { +class MIND_API Ceil : public BaseOperator { public: + MIND_API_BASE_MEMBER(Ceil); /// \brief Constructor. - Ceil() : PrimitiveC(kNameCeil) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Ceil() = default; - MS_DECLARE_PARENT(Ceil, PrimitiveC); + Ceil() : BaseOperator(kNameCeil) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Ceil for the inputs. void Init() const {} }; -AbstractBasePtr CeilInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CeilInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_CEIL_H_ diff --git a/mindspore/core/ops/celu.cc b/mindspore/core/ops/celu.cc index 74487a7a4c..91990187ac 100644 --- a/mindspore/core/ops/celu.cc +++ b/mindspore/core/ops/celu.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -48,6 +49,7 @@ TypePtr CeLUInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/celu.h b/mindspore/core/ops/celu.h index 18a7bf6b2c..afb9934b92 100644 --- a/mindspore/core/ops/celu.h +++ b/mindspore/core/ops/celu.h @@ -21,28 +21,25 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCeLU = "CeLU"; /// \brief Computes CeLU (Continuously differentiable exponential linear units) of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.CeLU for more details. -class CeLU : public PrimitiveC { +class CeLU : public BaseOperator { public: + MIND_API_BASE_MEMBER(CeLU); /// \brief Constructor. - CeLU() : PrimitiveC(kNameCeLU) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~CeLU() = default; - MS_DECLARE_PARENT(CeLU, PrimitiveC); + CeLU() : BaseOperator(kNameCeLU) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.CeLU for the inputs. void Init() {} }; -AbstractBasePtr CeLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CeLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCeLUPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cholesky_inverse_.cc b/mindspore/core/ops/cholesky_inverse_.cc index 8405df5c5e..b62e767bf2 100644 --- a/mindspore/core/ops/cholesky_inverse_.cc +++ b/mindspore/core/ops/cholesky_inverse_.cc @@ -17,6 +17,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -50,6 +51,7 @@ TypePtr CholeskyInverseInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/cholesky_inverse_.h b/mindspore/core/ops/cholesky_inverse_.h index 149b946a8b..647d3cf8f9 100644 --- a/mindspore/core/ops/cholesky_inverse_.h +++ b/mindspore/core/ops/cholesky_inverse_.h @@ -21,21 +21,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCholeskyInverse = "CholeskyInverse"; -class CholeskyInverse : public PrimitiveC { +class MIND_API CholeskyInverse : public BaseOperator { public: - CholeskyInverse() : PrimitiveC(kNameCholeskyInverse) { InitIOName({"x"}, {"y"}); } - ~CholeskyInverse() = default; - MS_DECLARE_PARENT(CholeskyInverse, PrimitiveC); + MIND_API_BASE_MEMBER(CholeskyInverse); + CholeskyInverse() : BaseOperator(kNameCholeskyInverse) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr CholeskyInverseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CholeskyInverseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCholeskyInversePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/clip.cc b/mindspore/core/ops/clip.cc index fb8d09acc7..e8e561f317 100644 --- a/mindspore/core/ops/clip.cc +++ b/mindspore/core/ops/clip.cc @@ -18,22 +18,24 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Clip, PrimitiveC, BaseOperator); void Clip::Init(const float max, const float min) { this->set_max(max); this->set_min(min); } -void Clip::set_max(const float max) { (void)this->AddAttr(kMax, MakeValue(max)); } +void Clip::set_max(const float max) { (void)this->AddAttr(kMax, api::MakeValue(max)); } float Clip::get_max() const { auto value_ptr = this->GetAttr(kMax); return GetValue(value_ptr); } -void Clip::set_min(const float min) { (void)this->AddAttr(kMin, MakeValue(min)); } +void Clip::set_min(const float min) { (void)this->AddAttr(kMin, api::MakeValue(min)); } float Clip::get_min() const { auto value_ptr = this->GetAttr(kMin); diff --git a/mindspore/core/ops/clip.h b/mindspore/core/ops/clip.h index ac20ad1f96..f2a989083e 100644 --- a/mindspore/core/ops/clip.h +++ b/mindspore/core/ops/clip.h @@ -17,23 +17,18 @@ #define MINDSPORE_CORE_OPS_CLIP_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameClip = "Clip"; /// \brief Clip defined Clip operator prototype. -class MS_CORE_API Clip : public PrimitiveC { +class MIND_API Clip : public BaseOperator { public: + MIND_API_BASE_MEMBER(Clip); /// \brief Constructor. - Clip() : PrimitiveC(kNameClip) {} - - /// \brief Destructor. - ~Clip() = default; - - MS_DECLARE_PARENT(Clip, PrimitiveC); + Clip() : BaseOperator(kNameClip) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/coalesce.cc b/mindspore/core/ops/coalesce.cc index 327cca879d..1f3aa9e6b6 100644 --- a/mindspore/core/ops/coalesce.cc +++ b/mindspore/core/ops/coalesce.cc @@ -18,6 +18,9 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" #include "utils/tensor_construct_utils.h" +#include "abstract/dshape.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -82,6 +85,8 @@ abstract::TupleShapePtr CoalesceInferShape(const PrimitivePtr &primitive, std::vector{y_indices_shape_list, y_values_shape_list, y_shape_shape_list}); } } // namespace + +MIND_API_BASE_IMPL(Coalesce, PrimitiveC, BaseOperator); AbstractBasePtr CoalesceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { const int64_t input_num = 3; diff --git a/mindspore/core/ops/coalesce.h b/mindspore/core/ops/coalesce.h index 2c2a077e12..ab18f48c12 100644 --- a/mindspore/core/ops/coalesce.h +++ b/mindspore/core/ops/coalesce.h @@ -22,25 +22,22 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "abstract/dshape.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCoalesce = "Coalesce"; -class Coalesce : public PrimitiveC { +class MIND_API Coalesce : public BaseOperator { public: - Coalesce() : PrimitiveC(kNameCoalesce) { + MIND_API_BASE_MEMBER(Coalesce); + Coalesce() : BaseOperator(kNameCoalesce) { InitIOName({"x_indices", "x_values", "x_shape"}, {"y_indices", "y_values", "y_shape"}); } - ~Coalesce() = default; - MS_DECLARE_PARENT(Coalesce, PrimitiveC); }; -AbstractBasePtr CoalesceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CoalesceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCoalescePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/complex.cc b/mindspore/core/ops/complex.cc index 8226605b4f..dbb055d24f 100644 --- a/mindspore/core/ops/complex.cc +++ b/mindspore/core/ops/complex.cc @@ -21,6 +21,8 @@ #include #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -127,6 +129,8 @@ ValuePtr ComplexInferValue(const PrimitivePtr &prim, const std::vector #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns a complex Tensor from the real and imaginary part. /// Refer to Python API @ref mindspore.ops.Complex for more details. -class MS_CORE_API Complex : public PrimitiveC { +class MIND_API Complex : public BaseOperator { public: + MIND_API_BASE_MEMBER(Complex); /// \brief Constructor. - Complex() : PrimitiveC(prim::kPrimSquare->name()) { InitIOName({"s", "input_imag"}, {"output"}); } - /// \brief Destructor. - ~Complex() = default; - MS_DECLARE_PARENT(Complex, PrimitiveC); + Complex() : BaseOperator("Square") { InitIOName({"s", "input_imag"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Complex for the inputs. void Init() {} }; diff --git a/mindspore/core/ops/concat.cc b/mindspore/core/ops/concat.cc index e2e3524126..c77e08fab6 100644 --- a/mindspore/core/ops/concat.cc +++ b/mindspore/core/ops/concat.cc @@ -19,15 +19,18 @@ #include "ops/concat.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Concat, PrimitiveC, BaseOperator); void Concat::Init(const int64_t axis) { this->set_axis(axis); } int64_t Concat::get_axis() const { auto value_ptr = this->GetAttr(kAxis); return GetValue(value_ptr); } -void Concat::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void Concat::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } REGISTER_PRIMITIVE_C(kNameConcat, Concat); } // namespace ops diff --git a/mindspore/core/ops/concat.h b/mindspore/core/ops/concat.h index a01639ea73..5da65e1201 100644 --- a/mindspore/core/ops/concat.h +++ b/mindspore/core/ops/concat.h @@ -19,23 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameConcat = "Concat"; /// \brief Connect tensor in the specified axis. /// Refer to Python API @ref mindspore.ops.Concat for more details. -class MS_CORE_API Concat : public PrimitiveC { +class MIND_API Concat : public BaseOperator { public: + MIND_API_BASE_MEMBER(Concat); /// \brief Constructor. - Concat() : PrimitiveC(kNameConcat) {} - /// \brief Destructor. - ~Concat() = default; - MS_DECLARE_PARENT(Concat, PrimitiveC); + Concat() : BaseOperator(kNameConcat) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Concat for the inputs. void Init(const int64_t axis = 0); /// \brief Set axis. @@ -45,8 +41,8 @@ class MS_CORE_API Concat : public PrimitiveC { /// \return axis. int64_t get_axis() const; }; -AbstractBasePtr ConcatInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ConcatInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_CONCAT_H_ diff --git a/mindspore/core/ops/conj.cc b/mindspore/core/ops/conj.cc index c3c8047225..2875a02f12 100644 --- a/mindspore/core/ops/conj.cc +++ b/mindspore/core/ops/conj.cc @@ -20,6 +20,8 @@ #include #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -47,6 +49,8 @@ AbstractBasePtr ConjInfer(const abstract::AnalysisEnginePtr &, const PrimitivePt return abstract::MakeAbstract(ConjInferShape(primitive, input_args), ConjInferType(primitive, input_args)); } } // namespace + +MIND_API_BASE_IMPL(Conj, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_EVAL_IMPL(Conj, prim::kPrimConj, ConjInfer, nullptr, true); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/conj.h b/mindspore/core/ops/conj.h index 830192585c..e6e24ad3f2 100644 --- a/mindspore/core/ops/conj.h +++ b/mindspore/core/ops/conj.h @@ -19,21 +19,18 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns a Tensor that is the conjugate part of the input. /// Refer to Python API @ref mindspore.ops.Conj for more details. -class MS_CORE_API Conj : public PrimitiveC { +class MIND_API Conj : public BaseOperator { public: + MIND_API_BASE_MEMBER(Conj); /// \brief Constructor. - Conj() : PrimitiveC(prim::kPrimConj->name()) { InitIOName({"input"}, {"output"}); } - /// \brief Destructor. - ~Conj() = default; - MS_DECLARE_PARENT(Conj, PrimitiveC); + Conj() : BaseOperator("Conj") { InitIOName({"input"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Conj for the inputs. void Init() {} }; diff --git a/mindspore/core/ops/constant_of_shape.cc b/mindspore/core/ops/constant_of_shape.cc index 7dc0d6bb65..f71296e46a 100644 --- a/mindspore/core/ops/constant_of_shape.cc +++ b/mindspore/core/ops/constant_of_shape.cc @@ -18,6 +18,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -26,20 +27,21 @@ void ConstantOfShape::Init(int64_t data_type, const std::vector &value) { this->set_value(value); } -void ConstantOfShape::set_data_type(int64_t data_type) { (void)this->AddAttr(kDataType, MakeValue(data_type)); } +void ConstantOfShape::set_data_type(int64_t data_type) { (void)this->AddAttr(kDataType, api::MakeValue(data_type)); } int64_t ConstantOfShape::get_data_type() const { auto value_ptr = this->GetAttr(kDataType); return GetValue(value_ptr); } -void ConstantOfShape::set_value(const std::vector &value) { (void)this->AddAttr(kValue, MakeValue(value)); } +void ConstantOfShape::set_value(const std::vector &value) { (void)this->AddAttr(kValue, api::MakeValue(value)); } std::vector ConstantOfShape::get_value() const { auto value_ptr = this->GetAttr(kValue); return GetValue>(value_ptr); } +MIND_API_BASE_IMPL(ConstantOfShape, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameConstantOfShape, ConstantOfShape); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/constant_of_shape.h b/mindspore/core/ops/constant_of_shape.h index af676ad7fd..97859b12b8 100644 --- a/mindspore/core/ops/constant_of_shape.h +++ b/mindspore/core/ops/constant_of_shape.h @@ -18,23 +18,18 @@ #define MINDSPORE_CORE_OPS_CONSTANT_OF_SHAPE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameConstantOfShape = "ConstantOfShape"; /// \brief ConstantOfShape defined ConstantOfShape operator prototype of lite. -class MS_CORE_API ConstantOfShape : public PrimitiveC { +class MIND_API ConstantOfShape : public BaseOperator { public: + MIND_API_BASE_MEMBER(ConstantOfShape); /// \brief Constructor. - ConstantOfShape() : PrimitiveC(kNameConstantOfShape) {} - - /// \brief Destructor. - ~ConstantOfShape() = default; - - MS_DECLARE_PARENT(ConstantOfShape, PrimitiveC); + ConstantOfShape() : BaseOperator(kNameConstantOfShape) {} /// \brief Method to init the op's attributes. /// @@ -63,8 +58,8 @@ class MS_CORE_API ConstantOfShape : public PrimitiveC { std::vector get_value() const; }; -AbstractBasePtr ConstantOfShapeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ConstantOfShapeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/control_depend.cc b/mindspore/core/ops/control_depend.cc index 97822996a8..df4d67581d 100644 --- a/mindspore/core/ops/control_depend.cc +++ b/mindspore/core/ops/control_depend.cc @@ -15,14 +15,18 @@ */ #include "ops/control_depend.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ControlDepend, PrimitiveC, BaseOperator); void ControlDepend::Init(const int64_t depend_mode) { this->set_depend_mode(depend_mode); } void ControlDepend::set_depend_mode(const int64_t depend_mode) { CheckAndConvertUtils::CheckInRange(kDependMode, depend_mode, kIncludeBoth, {0, 1}, name()); - (void)AddAttr(kDependMode, MakeValue(depend_mode)); + (void)AddAttr(kDependMode, api::MakeValue(depend_mode)); } int64_t ControlDepend::get_depend_mode() const { diff --git a/mindspore/core/ops/control_depend.h b/mindspore/core/ops/control_depend.h index f6cd375506..bbb4b2185a 100644 --- a/mindspore/core/ops/control_depend.h +++ b/mindspore/core/ops/control_depend.h @@ -19,19 +19,16 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameControlDepend = "ControlDepend"; -class MS_CORE_API ControlDepend : public PrimitiveC { +class MIND_API ControlDepend : public BaseOperator { public: - ControlDepend() : PrimitiveC(kNameControlDepend) {} - ~ControlDepend() = default; - MS_DECLARE_PARENT(ControlDepend, PrimitiveC); + MIND_API_BASE_MEMBER(ControlDepend); + ControlDepend() : BaseOperator(kNameControlDepend) {} void Init(const int64_t depend_mode); void set_depend_mode(const int64_t depend_mode = 0); int64_t get_depend_mode() const; diff --git a/mindspore/core/ops/conv2d.cc b/mindspore/core/ops/conv2d.cc index 1e02cd3359..2b53c43541 100644 --- a/mindspore/core/ops/conv2d.cc +++ b/mindspore/core/ops/conv2d.cc @@ -23,6 +23,8 @@ #include "ir/dtype/tensor_type.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" using mindspore::abstract::Shape; namespace mindspore { @@ -254,6 +256,8 @@ TypePtr Conv2dInferType(const PrimitivePtr &prim, const std::vector &kernel_size, int64_t mode, const PadMode &pad_mode, const std::vector &pad, const std::vector &stride, const std::vector &dilation, int64_t group, const Format &format) { @@ -270,19 +274,20 @@ void Conv2D::Init(int64_t out_channel, const std::vector &kernel_size, void Conv2D::set_out_channel(int64_t out_channel) { (void)AddAttr(kOutChannel, - MakeValue(CheckAndConvertUtils::CheckInteger(kOutChannel, out_channel, kGreaterThan, 0, name()))); + api::MakeValue(CheckAndConvertUtils::CheckInteger(kOutChannel, out_channel, kGreaterThan, 0, name()))); } void Conv2D::set_kernel_size(const std::vector &kernel_size) { - (void)AddAttr(kKernelSize, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, name()))); + (void)AddAttr(kKernelSize, + api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, name()))); } void Conv2D::set_stride(const std::vector &stride) { - (void)AddAttr(kStride, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStride, stride, name()))); + (void)AddAttr(kStride, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStride, stride, name()))); } void Conv2D::set_dilation(const std::vector &dilation) { - (void)AddAttr(kDilation, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kDilation, dilation, name()))); + (void)AddAttr(kDilation, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kDilation, dilation, name()))); } void Conv2D::set_pad_mode(const PadMode &pad_mode) { @@ -295,26 +300,26 @@ void Conv2D::set_pad_mode(const PadMode &pad_mode) { CheckAndConvertUtils::Check(kPad, pad, kEqual, {0, 0, 0, 0}, name()); } int64_t swi = pad_mode; - (void)AddAttr(kPadMode, MakeValue(swi)); + (void)AddAttr(kPadMode, api::MakeValue(swi)); } void Conv2D::set_pad(const std::vector &pad) { const int64_t pad_size = 4; (void)CheckAndConvertUtils::CheckInteger("pad_size", SizeToLong(pad.size()), kEqual, pad_size, name()); - (void)AddAttr(kPad, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kPad, pad, name()))); + (void)AddAttr(kPad, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kPad, pad, name()))); } void Conv2D::set_mode(int64_t mode) { - (void)AddAttr(kMode, MakeValue(CheckAndConvertUtils::CheckInteger(kMode, mode, kEqual, 1, name()))); + (void)AddAttr(kMode, api::MakeValue(CheckAndConvertUtils::CheckInteger(kMode, mode, kEqual, 1, name()))); } void Conv2D::set_group(int64_t group) { - (void)AddAttr(kGroup, MakeValue(CheckAndConvertUtils::CheckInteger(kGroup, group, kGreaterThan, 0, name()))); + (void)AddAttr(kGroup, api::MakeValue(CheckAndConvertUtils::CheckInteger(kGroup, group, kGreaterThan, 0, name()))); } void Conv2D::set_format(const Format &format) { int64_t f = format; - (void)AddAttr(kFormat, MakeValue(f)); + (void)AddAttr(kFormat, api::MakeValue(f)); } int64_t Conv2D::get_out_channel() const { diff --git a/mindspore/core/ops/conv2d.h b/mindspore/core/ops/conv2d.h index 3d1ad360cd..7ae9927493 100644 --- a/mindspore/core/ops/conv2d.h +++ b/mindspore/core/ops/conv2d.h @@ -21,22 +21,20 @@ #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" + namespace mindspore { namespace ops { constexpr auto kNameConv2D = "Conv2D"; /// \brief 2D convolution layer. Refer to Python API @ref mindspore.ops.Conv2D for more details. -class MS_CORE_API Conv2D : public PrimitiveC { +class MIND_API Conv2D : public BaseOperator { public: + MIND_API_BASE_MEMBER(Conv2D); /// \brief Constructor. - Conv2D() : PrimitiveC(kNameConv2D) { InitIOName({"x", "w"}, {"output"}); } - explicit Conv2D(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "w"}, {"output"}); } - /// \brief Destructor. - ~Conv2D() = default; - MS_DECLARE_PARENT(Conv2D, PrimitiveC); + Conv2D() : BaseOperator(kNameConv2D) { InitIOName({"x", "w"}, {"output"}); } + explicit Conv2D(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "w"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Conv2D for the inputs. void Init(int64_t out_channel, const std::vector &kernel_size, int64_t mode = 1, const PadMode &pad_mode = VALID, const std::vector &pad = {0, 0, 0, 0}, @@ -97,8 +95,8 @@ class MS_CORE_API Conv2D : public PrimitiveC { /// \return format. Format get_format() const; }; -AbstractBasePtr Conv2dInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr Conv2dInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_CONV2D_H_ diff --git a/mindspore/core/ops/conv2d_transpose.cc b/mindspore/core/ops/conv2d_transpose.cc index 7248e17634..a9a57ef508 100644 --- a/mindspore/core/ops/conv2d_transpose.cc +++ b/mindspore/core/ops/conv2d_transpose.cc @@ -21,10 +21,14 @@ #include #include "ops/conv2d_transpose.h" +#include "ops/op_utils.h" #include "ops/grad/conv2d_backprop_input.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Conv2DTranspose, PrimitiveC, BaseOperator); void Conv2DTranspose::Init(int64_t in_channel, int64_t out_channel, const std::vector &kernel_size, int64_t mode, const PadMode &pad_mode, const std::vector &pad, const std::vector &stride, const std::vector &dilation, int64_t group, @@ -44,12 +48,12 @@ void Conv2DTranspose::Init(int64_t in_channel, int64_t out_channel, const std::v void Conv2DTranspose::set_in_channel(int64_t in_channel) { (void)AddAttr(kInChannel, - MakeValue(CheckAndConvertUtils::CheckInteger(kInChannel, in_channel, kGreaterThan, 0, name()))); + api::MakeValue(CheckAndConvertUtils::CheckInteger(kInChannel, in_channel, kGreaterThan, 0, name()))); } void Conv2DTranspose::set_out_channel(int64_t out_channel) { (void)AddAttr(kOutChannel, - MakeValue(CheckAndConvertUtils::CheckInteger(kOutChannel, out_channel, kGreaterThan, 0, name()))); + api::MakeValue(CheckAndConvertUtils::CheckInteger(kOutChannel, out_channel, kGreaterThan, 0, name()))); } void Conv2DTranspose::set_kernel_size(const std::vector &kernel_size) { @@ -58,7 +62,7 @@ void Conv2DTranspose::set_kernel_size(const std::vector &kernel_size) { for (int64_t item : kernel_size) { (void)CheckAndConvertUtils::CheckInteger(kKernelSize, item, kGreaterEqual, 1, name()); } - (void)AddAttr(kKernelSize, MakeValue(kernel_size)); + (void)AddAttr(kKernelSize, api::MakeValue(kernel_size)); } void Conv2DTranspose::set_stride(const std::vector &stride) { @@ -67,14 +71,14 @@ void Conv2DTranspose::set_stride(const std::vector &stride) { for (int64_t item : stride) { (void)CheckAndConvertUtils::CheckInteger(kStride, item, kGreaterEqual, 1, name()); } - (void)AddAttr(kStride, MakeValue(stride)); + (void)AddAttr(kStride, api::MakeValue(stride)); } void Conv2DTranspose::set_dilation(const std::vector &dilation) { const int64_t dilation_size = 2; (void)CheckAndConvertUtils::CheckInteger(kDilation, SizeToLong(dilation.size()), kGreaterEqual, dilation_size, name()); - (void)AddAttr(kDilation, MakeValue(dilation)); + (void)AddAttr(kDilation, api::MakeValue(dilation)); } void Conv2DTranspose::set_pad_mode(const PadMode &pad_mode) { @@ -87,32 +91,32 @@ void Conv2DTranspose::set_pad_mode(const PadMode &pad_mode) { CheckAndConvertUtils::Check(kPad, pad, kEqual, {0, 0, 0, 0}, name()); } int64_t swi = pad_mode; - (void)AddAttr(kPadMode, MakeValue(swi)); + (void)AddAttr(kPadMode, api::MakeValue(swi)); } void Conv2DTranspose::set_pad(const std::vector &pad) { const int64_t pad_size = 4; (void)CheckAndConvertUtils::CheckInteger("pad_size", SizeToLong(pad.size()), kEqual, pad_size, name()); - (void)AddAttr(kPad, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kPad, pad, name()))); + (void)AddAttr(kPad, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kPad, pad, name()))); } void Conv2DTranspose::set_mode(int64_t mode) { - (void)AddAttr(kMode, MakeValue(CheckAndConvertUtils::CheckInteger(kMode, mode, kEqual, 1, name()))); + (void)AddAttr(kMode, api::MakeValue(CheckAndConvertUtils::CheckInteger(kMode, mode, kEqual, 1, name()))); } void Conv2DTranspose::set_group(int64_t group) { - (void)AddAttr(kGroup, MakeValue(CheckAndConvertUtils::CheckInteger(kGroup, group, kGreaterThan, 0, name()))); + (void)AddAttr(kGroup, api::MakeValue(CheckAndConvertUtils::CheckInteger(kGroup, group, kGreaterThan, 0, name()))); } void Conv2DTranspose::set_format(const Format &format) { int64_t f = format; - (void)AddAttr(kFormat, MakeValue(f)); + (void)AddAttr(kFormat, api::MakeValue(f)); } void Conv2DTranspose::set_pad_list(const std::vector &pad_list) { const int64_t pad_size = 4; (void)CheckAndConvertUtils::CheckInteger(kPadList, SizeToLong(pad_list.size()), kEqual, pad_size, name()); - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } int64_t Conv2DTranspose::get_in_channel() const { diff --git a/mindspore/core/ops/conv2d_transpose.h b/mindspore/core/ops/conv2d_transpose.h index 0e80a59721..6129bd2850 100644 --- a/mindspore/core/ops/conv2d_transpose.h +++ b/mindspore/core/ops/conv2d_transpose.h @@ -21,26 +21,24 @@ #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" + namespace mindspore { namespace ops { constexpr auto kNameConv2DTranspose = "Conv2DTranspose"; /// \brief 2D transposed convolution layer. Refer to Python API @ref mindspore.nn.Conv2dTranspose for more details. -class MS_CORE_API Conv2DTranspose : public PrimitiveC { +class MIND_API Conv2DTranspose : public BaseOperator { public: + MIND_API_BASE_MEMBER(Conv2DTranspose); /// \brief Constructor. - Conv2DTranspose() : PrimitiveC(kNameConv2DTranspose) { + Conv2DTranspose() : BaseOperator(kNameConv2DTranspose) { InitIOName({"out_backprop", "filter", "input_sizes"}, {"output"}); } - explicit Conv2DTranspose(const std::string k_name) : PrimitiveC(k_name) { + explicit Conv2DTranspose(const std::string k_name) : BaseOperator(k_name) { InitIOName({"out_backprop", "filter", "input_sizes"}, {"output"}); } - /// \brief Destructor. - ~Conv2DTranspose() = default; - MS_DECLARE_PARENT(Conv2DTranspose, PrimitiveC); /// \brief Init. Refer to the parameters of Python API @ref mindspore.nn.Conv2dTranspose for the inputs. void Init(int64_t in_channel, int64_t out_channel, const std::vector &kernel_size, int64_t mode = 1, const PadMode &pad_mode = VALID, const std::vector &pad = {0, 0, 0, 0}, @@ -114,8 +112,8 @@ class MS_CORE_API Conv2DTranspose : public PrimitiveC { /// \return pad_list. std::vector get_pad_list() const; }; -AbstractBasePtr Conv2DTransposeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr Conv2DTransposeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_CONV2D_TRANSPOSE_H_ diff --git a/mindspore/core/ops/cos.cc b/mindspore/core/ops/cos.cc index b8187aa2c1..dd720ebbf0 100644 --- a/mindspore/core/ops/cos.cc +++ b/mindspore/core/ops/cos.cc @@ -18,6 +18,9 @@ #include #include #include +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +44,7 @@ TypePtr CosInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/cos.h b/mindspore/core/ops/cos.h index 28dbc671f4..6a7bb33d08 100644 --- a/mindspore/core/ops/cos.h +++ b/mindspore/core/ops/cos.h @@ -19,26 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { /// \brief Computes cosine of input element-wise. Refer to Python API @ref mindspore.ops.Cos for more details. -class MS_CORE_API Cos : public PrimitiveC { +class MIND_API Cos : public BaseOperator { public: + MIND_API_BASE_MEMBER(Cos); /// \brief Constructor. - Cos() : PrimitiveC(prim::kPrimCos->name()) {} - /// \brief Destructor. - ~Cos() = default; - MS_DECLARE_PARENT(Cos, PrimitiveC); + Cos() : BaseOperator("Cos") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Cos for the inputs. void Init(float alpha = 0.0); }; -AbstractBasePtr CosInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CosInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_COS_H_ diff --git a/mindspore/core/ops/cosh.cc b/mindspore/core/ops/cosh.cc index afbb482f55..b9ce63f0bf 100644 --- a/mindspore/core/ops/cosh.cc +++ b/mindspore/core/ops/cosh.cc @@ -18,6 +18,11 @@ #include #include #include +#include +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +50,7 @@ TypePtr CoshInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/cosh.h b/mindspore/core/ops/cosh.h index 6a5af9c21a..85cf0f830e 100644 --- a/mindspore/core/ops/cosh.h +++ b/mindspore/core/ops/cosh.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_COSH_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCosh = "Cosh"; -class Cosh : public PrimitiveC { +class Cosh : public BaseOperator { public: - Cosh() : PrimitiveC(kNameCosh) { InitIOName({"x"}, {"output"}); } - ~Cosh() = default; - MS_DECLARE_PARENT(Cosh, PrimitiveC); + MIND_API_BASE_MEMBER(Cosh); + Cosh() : BaseOperator(kNameCosh) { InitIOName({"x"}, {"output"}); } void Init() {} }; -AbstractBasePtr CoshInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CoshInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimCoshPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/crop.cc b/mindspore/core/ops/crop.cc index 123ca821c0..ebd212bea2 100644 --- a/mindspore/core/ops/crop.cc +++ b/mindspore/core/ops/crop.cc @@ -19,22 +19,24 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Crop, PrimitiveC, BaseOperator); void Crop::Init(const int64_t axis, const std::vector &offsets) { this->set_axis(axis); this->set_offsets(offsets); } -void Crop::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void Crop::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t Crop::get_axis() const { auto value_ptr = this->GetAttr(kAxis); return GetValue(value_ptr); } -void Crop::set_offsets(const std::vector &offsets) { (void)this->AddAttr(kOffsets, MakeValue(offsets)); } +void Crop::set_offsets(const std::vector &offsets) { (void)this->AddAttr(kOffsets, api::MakeValue(offsets)); } std::vector Crop::get_offsets() const { auto value_ptr = this->GetAttr(kOffsets); diff --git a/mindspore/core/ops/crop.h b/mindspore/core/ops/crop.h index ce31801cc3..89212dfafc 100644 --- a/mindspore/core/ops/crop.h +++ b/mindspore/core/ops/crop.h @@ -19,23 +19,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCrop = "Crop"; /// \brief Crop defined the Crop operator prototype of lite, which can be replaced by slice operator. -class MS_CORE_API Crop : public PrimitiveC { +class MIND_API Crop : public BaseOperator { public: + MIND_API_BASE_MEMBER(Crop); /// \brief Constructor. - Crop() : PrimitiveC(kNameCrop) {} - - /// \brief Destructor. - ~Crop() = default; - - MS_DECLARE_PARENT(Crop, PrimitiveC); + Crop() : BaseOperator(kNameCrop) {} /// \brief Method to init the op's attributes. /// @@ -63,8 +58,8 @@ class MS_CORE_API Crop : public PrimitiveC { /// \return a vector which indicates the start index to slice on the corresponding axis. std::vector get_offsets() const; }; -AbstractBasePtr CropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCrop = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/crop_and_resize.cc b/mindspore/core/ops/crop_and_resize.cc index b1e6ab9010..71eb98683b 100644 --- a/mindspore/core/ops/crop_and_resize.cc +++ b/mindspore/core/ops/crop_and_resize.cc @@ -19,9 +19,11 @@ #include "ops/crop_and_resize.h" #include "utils/check_convert_utils.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(CropAndResize, PrimitiveC, BaseOperator); void CropAndResize::Init(ResizeMethod method, float extrapolation_value) { this->set_method(method); this->set_extrapolation_value(extrapolation_value); @@ -29,11 +31,11 @@ void CropAndResize::Init(ResizeMethod method, float extrapolation_value) { void CropAndResize::set_method(ResizeMethod method) { auto swi = (int64_t)method; - (void)this->AddAttr(kMethod, MakeValue(swi)); + (void)this->AddAttr(kMethod, api::MakeValue(swi)); } void CropAndResize::set_extrapolation_value(float extrapolation_value) { - (void)this->AddAttr(kExtrapolationValue, MakeValue(extrapolation_value)); + (void)this->AddAttr(kExtrapolationValue, api::MakeValue(extrapolation_value)); } ResizeMethod CropAndResize::get_method() const { diff --git a/mindspore/core/ops/crop_and_resize.h b/mindspore/core/ops/crop_and_resize.h index 6d3d082d2f..4571f54234 100644 --- a/mindspore/core/ops/crop_and_resize.h +++ b/mindspore/core/ops/crop_and_resize.h @@ -18,22 +18,19 @@ #define MINDSPORE_CORE_OPS_CROP_AND_RESIZE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCropAndResize = "CropAndResize"; /// \brief Extracts crops from the input image tensor and resizes them. /// Refer to Python API @ref mindspore.ops.CropAndResize for more details. -class MS_CORE_API CropAndResize : public PrimitiveC { +class MIND_API CropAndResize : public BaseOperator { public: + MIND_API_BASE_MEMBER(CropAndResize); /// \brief Constructor. - CropAndResize() : PrimitiveC(kNameCropAndResize) { InitIOName({"x", "boxes", "box_index", "crop_size"}, {"y"}); } - /// \brief Destructor. - ~CropAndResize() = default; - MS_DECLARE_PARENT(CropAndResize, PrimitiveC); + CropAndResize() : BaseOperator(kNameCropAndResize) { InitIOName({"x", "boxes", "box_index", "crop_size"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.CropAndResize for the inputs. void Init(ResizeMethod method, float extrapolation_value); diff --git a/mindspore/core/ops/cross.cc b/mindspore/core/ops/cross.cc index 22a41feb6d..ef4836e9d2 100644 --- a/mindspore/core/ops/cross.cc +++ b/mindspore/core/ops/cross.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -102,5 +103,7 @@ AbstractBasePtr CrossInfer(const abstract::AnalysisEnginePtr &, const PrimitiveP REGISTER_PRIMITIVE_EVAL_IMPL(Cross, prim::kPrimCross, CrossInfer, nullptr, true); } // namespace + +MIND_API_BASE_IMPL(Cross, PrimitiveC, BaseOperator); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cross.h b/mindspore/core/ops/cross.h index dee7902ad5..b38534cf18 100644 --- a/mindspore/core/ops/cross.h +++ b/mindspore/core/ops/cross.h @@ -20,23 +20,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCross = "Cross"; -class Cross : public PrimitiveC { +class MIND_API Cross : public BaseOperator { public: - Cross() : PrimitiveC(kNameCross) { InitIOName({"x1", "x2"}, {"y"}); } - ~Cross() = default; - MS_DECLARE_PARENT(Cross, PrimitiveC); + MIND_API_BASE_MEMBER(Cross); + Cross() : BaseOperator(kNameCross) { InitIOName({"x1", "x2"}, {"y"}); } }; -AbstractBasePtr CrossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CrossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCrossPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/ctc_loss_v2.cc b/mindspore/core/ops/ctc_loss_v2.cc index c949928669..5c8da45dad 100644 --- a/mindspore/core/ops/ctc_loss_v2.cc +++ b/mindspore/core/ops/ctc_loss_v2.cc @@ -24,6 +24,8 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -71,6 +73,7 @@ TuplePtr CTCLossV2InferType(const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(CTCLossV2InferShape(primitive, input_args), CTCLossV2InferType(primitive, input_args)); diff --git a/mindspore/core/ops/ctc_loss_v2.h b/mindspore/core/ops/ctc_loss_v2.h index a7fa96c97f..2d62101a45 100644 --- a/mindspore/core/ops/ctc_loss_v2.h +++ b/mindspore/core/ops/ctc_loss_v2.h @@ -21,29 +21,27 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" + namespace mindspore { namespace ops { constexpr auto kNameCTCLossV2 = "CTCLossV2"; /// \brief Calculates the CTC (Connectionist Temporal Classification) loss and the gradient. /// Refer to Python API @ref mindspore.ops.CTCLossV2 for more details. -class MS_CORE_API CTCLossV2 : public PrimitiveC { +class MIND_API CTCLossV2 : public BaseOperator { public: + MIND_API_BASE_MEMBER(CTCLossV2); /// \brief Constructor. - CTCLossV2() : PrimitiveC(kNameCTCLossV2) { + CTCLossV2() : BaseOperator(kNameCTCLossV2) { InitIOName({"log_probs", "targets", "input_lengths", "target_lengths"}, {"neg_log_likelihood", "log_alpha"}); } - /// \brief Destructor. - ~CTCLossV2() = default; - MS_DECLARE_PARENT(CTCLossV2, PrimitiveC); /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.CTCLossV2 for the inputs. void Init() const {} }; -AbstractBasePtr CTCLossV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CTCLossV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/ctc_loss_v2_grad.cc b/mindspore/core/ops/ctc_loss_v2_grad.cc index 74bd2e2355..3ed1ff5a92 100644 --- a/mindspore/core/ops/ctc_loss_v2_grad.cc +++ b/mindspore/core/ops/ctc_loss_v2_grad.cc @@ -23,6 +23,8 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -65,6 +67,7 @@ TypePtr CTCLossV2GradInferType(const PrimitivePtr &primitive, const std::vector< } } // namespace +MIND_API_BASE_IMPL(CTCLossV2Grad, PrimitiveC, BaseOperator); AbstractBasePtr CTCLossV2GradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_shape = CTCLossV2GradInferShape(primitive, input_args); diff --git a/mindspore/core/ops/ctc_loss_v2_grad.h b/mindspore/core/ops/ctc_loss_v2_grad.h index 49e51443e2..feee20f66f 100644 --- a/mindspore/core/ops/ctc_loss_v2_grad.h +++ b/mindspore/core/ops/ctc_loss_v2_grad.h @@ -21,26 +21,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" + namespace mindspore { namespace ops { constexpr auto kNameCTCLossV2Grad = "CTCLossV2Grad"; -class MS_CORE_API CTCLossV2Grad : public PrimitiveC { +class MIND_API CTCLossV2Grad : public BaseOperator { public: - CTCLossV2Grad() : PrimitiveC(kNameCTCLossV2Grad) { + MIND_API_BASE_MEMBER(CTCLossV2Grad); + CTCLossV2Grad() : BaseOperator(kNameCTCLossV2Grad) { InitIOName( {"grad_out", "log_probs", "targets", "input_lengths", "target_lengths", "neg_log_likelihood", "log_alpha"}, {"grad"}); } - ~CTCLossV2Grad() = default; - MS_DECLARE_PARENT(CTCLossV2Grad, PrimitiveC); void Init() const {} }; -AbstractBasePtr CTCLossV2GradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CTCLossV2GradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCTCLossV2Ptr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/ctcloss.cc b/mindspore/core/ops/ctcloss.cc index 45447e2a71..23ea65d1eb 100644 --- a/mindspore/core/ops/ctcloss.cc +++ b/mindspore/core/ops/ctcloss.cc @@ -24,6 +24,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -106,6 +107,7 @@ TuplePtr CTCLossInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/ctcloss.h b/mindspore/core/ops/ctcloss.h index 1eca1b38b7..4f2003a417 100644 --- a/mindspore/core/ops/ctcloss.h +++ b/mindspore/core/ops/ctcloss.h @@ -18,26 +18,23 @@ #define MINDSPORE_CORE_OPS_CTCLOSS_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Calculates the CTC (Connectionist Temporal Classification) loss and the gradient. /// Refer to Python API @ref mindspore.ops.CTCLoss for more details. -class MS_CORE_API CTCLoss : public PrimitiveC { +class MIND_API CTCLoss : public BaseOperator { public: + MIND_API_BASE_MEMBER(CTCLoss); /// \brief Constructor. - CTCLoss() : PrimitiveC(prim::kPrimCTCLoss->name()) {} - /// \brief Destructor. - ~CTCLoss() = default; - MS_DECLARE_PARENT(CTCLoss, PrimitiveC); + CTCLoss() : BaseOperator("CTCLoss") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.CTCLoss for the inputs. void Init() const {} }; -AbstractBasePtr CTCLossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CTCLossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cummax.cc b/mindspore/core/ops/cummax.cc index 50dcedb3c7..7250a5e439 100644 --- a/mindspore/core/ops/cummax.cc +++ b/mindspore/core/ops/cummax.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -62,5 +63,7 @@ AbstractBasePtr CummaxInfer(const abstract::AnalysisEnginePtr &, const Primitive REGISTER_PRIMITIVE_EVAL_IMPL(Cummax, prim::kPrimCummax, CummaxInfer, nullptr, true); } // namespace + +MIND_API_BASE_IMPL(Cummax, PrimitiveC, BaseOperator); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cummax.h b/mindspore/core/ops/cummax.h index 1e59f6d68b..1a87ebb874 100644 --- a/mindspore/core/ops/cummax.h +++ b/mindspore/core/ops/cummax.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCummax = "Cummax"; -class Cummax : public PrimitiveC { +class MIND_API Cummax : public BaseOperator { public: - Cummax() : PrimitiveC(kNameCummax) { InitIOName({"x"}, {"y", "indices"}); } - ~Cummax() = default; - MS_DECLARE_PARENT(Cummax, PrimitiveC); + MIND_API_BASE_MEMBER(Cummax); + Cummax() : BaseOperator(kNameCummax) { InitIOName({"x"}, {"y", "indices"}); } }; -AbstractBasePtr CummaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CummaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCummaxPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cummin.cc b/mindspore/core/ops/cummin.cc index a0cb22cb3c..02bb55205a 100644 --- a/mindspore/core/ops/cummin.cc +++ b/mindspore/core/ops/cummin.cc @@ -19,6 +19,7 @@ #include "abstract/primitive_infer_map.h" #include "utils/check_convert_utils.h" #include "ops/cummin.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -50,6 +51,7 @@ TuplePtr CumminInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto prim_name = primitive->name(); diff --git a/mindspore/core/ops/cummin.h b/mindspore/core/ops/cummin.h index c5ed68d8cb..2cfd00684f 100644 --- a/mindspore/core/ops/cummin.h +++ b/mindspore/core/ops/cummin.h @@ -21,23 +21,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCummin = "Cummin"; -class Cummin : public PrimitiveC { +class Cummin : public BaseOperator { public: - Cummin() : PrimitiveC(kNameCummin) { InitIOName({"x"}, {"y"}); } - ~Cummin() = default; - MS_DECLARE_PARENT(Cummin, PrimitiveC); + MIND_API_BASE_MEMBER(Cummin); + Cummin() : BaseOperator(kNameCummin) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr CumminInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CumminInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCumminPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/cumsum.cc b/mindspore/core/ops/cumsum.cc index 9ea4dd17d7..bef83886e5 100644 --- a/mindspore/core/ops/cumsum.cc +++ b/mindspore/core/ops/cumsum.cc @@ -19,22 +19,24 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(CumSum, PrimitiveC, BaseOperator); void CumSum::Init(const bool exclusive, const bool reverse) { this->set_exclusive(exclusive); this->set_reverse(reverse); } -void CumSum::set_exclusive(const bool exclusive) { (void)this->AddAttr(kExclusive, MakeValue(exclusive)); } +void CumSum::set_exclusive(const bool exclusive) { (void)this->AddAttr(kExclusive, api::MakeValue(exclusive)); } bool CumSum::get_exclusive() const { auto value_ptr = this->GetAttr(kExclusive); return GetValue(value_ptr); } -void CumSum::set_reverse(const bool reverse) { (void)this->AddAttr(kReverse, MakeValue(reverse)); } +void CumSum::set_reverse(const bool reverse) { (void)this->AddAttr(kReverse, api::MakeValue(reverse)); } bool CumSum::get_reverse() const { auto value_ptr = this->GetAttr(kReverse); diff --git a/mindspore/core/ops/cumsum.h b/mindspore/core/ops/cumsum.h index 51daee23aa..5e60d288c6 100644 --- a/mindspore/core/ops/cumsum.h +++ b/mindspore/core/ops/cumsum.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCumSum = "CumSum"; /// \brief Computes the cumulative sum of input tensor along axis. /// Refer to Python API @ref mindspore.ops.CumSum for more details. -class MS_CORE_API CumSum : public PrimitiveC { +class MIND_API CumSum : public BaseOperator { public: + MIND_API_BASE_MEMBER(CumSum); /// \brief Constructor. - CumSum() : PrimitiveC(kNameCumSum) {} - /// \brief Destructor. - ~CumSum() = default; - MS_DECLARE_PARENT(CumSum, PrimitiveC); + CumSum() : BaseOperator(kNameCumSum) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.CumSum for the inputs. void Init(const bool exclusive, const bool reverse); /// \brief Set exclusive. @@ -50,8 +47,8 @@ class MS_CORE_API CumSum : public PrimitiveC { /// \return reverse. bool get_reverse() const; }; -AbstractBasePtr CumSumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CumSumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimCumSum = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/custom.cc b/mindspore/core/ops/custom.cc index 4b1e75ed50..a7912e0a9d 100644 --- a/mindspore/core/ops/custom.cc +++ b/mindspore/core/ops/custom.cc @@ -15,15 +15,20 @@ */ #include "ops/custom.h" +#include +#include +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Custom, PrimitiveC, BaseOperator); void Custom::Init(const std::string &type, const std::map> &attrs) { this->set_type(type); this->set_attr(attrs); } -void Custom::set_type(const std::string &type) { (void)this->AddAttr(kType, MakeValue(type)); } +void Custom::set_type(const std::string &type) { (void)this->AddAttr(kType, api::MakeValue(type)); } std::string Custom::get_type() const { auto value_ptr = this->GetAttr(kType); @@ -31,17 +36,17 @@ std::string Custom::get_type() const { } void Custom::set_attr(const std::map> &attrs) { - ValuePtrList value_ptr_list; + api::ValuePtrList value_ptr_list; for (const auto &attr : attrs) { - (void)value_ptr_list.emplace_back(MakeValue(attr.first)); - (void)value_ptr_list.emplace_back(MakeValue>(attr.second)); + (void)value_ptr_list.emplace_back(api::MakeValue(attr.first)); + (void)value_ptr_list.emplace_back(api::MakeValue>(attr.second)); } - (void)this->AddAttr(kAttr, MakeValue(value_ptr_list)); + (void)this->AddAttr(kAttr, api::MakeValue(value_ptr_list)); } std::map> Custom::get_attr() const { std::map> attrs; - auto value_ptr_list = GetValue(this->GetAttr(kAttr)); + auto value_ptr_list = GetValue(this->GetAttr(kAttr)); for (size_t i = 0; i < value_ptr_list.size(); i += 2) { auto key = GetValue(value_ptr_list[i]); auto value = GetValue>(value_ptr_list[i + 1]); diff --git a/mindspore/core/ops/custom.h b/mindspore/core/ops/custom.h index b27adfa5f1..af3bd41f56 100644 --- a/mindspore/core/ops/custom.h +++ b/mindspore/core/ops/custom.h @@ -21,23 +21,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "ir/anf.h" + +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameCustom = "Custom"; /// \brief Custom defined user-defined operator prototype. -class MS_CORE_API Custom : public PrimitiveC { +class MIND_API Custom : public BaseOperator { public: + MIND_API_BASE_MEMBER(Custom); /// \brief Constructor. - Custom() : PrimitiveC(kNameCustom) {} - - /// \brief Destructor. - ~Custom() override = default; - - MS_DECLARE_PARENT(Custom, PrimitiveC); + Custom() : BaseOperator(kNameCustom) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/custom_extract_features.cc b/mindspore/core/ops/custom_extract_features.cc index ada6ea6ff1..f14328f630 100644 --- a/mindspore/core/ops/custom_extract_features.cc +++ b/mindspore/core/ops/custom_extract_features.cc @@ -18,9 +18,11 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(CustomExtractFeatures, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameCustomExtractFeatures, CustomExtractFeatures); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/custom_extract_features.h b/mindspore/core/ops/custom_extract_features.h index ba4c42a4a6..0111170c30 100644 --- a/mindspore/core/ops/custom_extract_features.h +++ b/mindspore/core/ops/custom_extract_features.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_CUSTOM_EXTRACT_FEATURES_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCustomExtractFeatures = "CustomExtractFeatures"; -class MS_CORE_API CustomExtractFeatures : public PrimitiveC { +class MIND_API CustomExtractFeatures : public BaseOperator { public: - CustomExtractFeatures() : PrimitiveC(kNameCustomExtractFeatures) {} - ~CustomExtractFeatures() = default; - MS_DECLARE_PARENT(CustomExtractFeatures, PrimitiveC); + MIND_API_BASE_MEMBER(CustomExtractFeatures); + CustomExtractFeatures() : BaseOperator(kNameCustomExtractFeatures) {} void Init() const {} }; -AbstractBasePtr CustomExtractFeaturesInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CustomExtractFeaturesInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/custom_normalize.cc b/mindspore/core/ops/custom_normalize.cc index eb89a3eee8..4956fe8f2d 100644 --- a/mindspore/core/ops/custom_normalize.cc +++ b/mindspore/core/ops/custom_normalize.cc @@ -17,9 +17,11 @@ #include "ops/custom_normalize.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(CustomNormalize, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameCustomNormalize, CustomNormalize); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/custom_normalize.h b/mindspore/core/ops/custom_normalize.h index 0377f39c69..f3f59f3f2d 100644 --- a/mindspore/core/ops/custom_normalize.h +++ b/mindspore/core/ops/custom_normalize.h @@ -18,23 +18,21 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCustomNormalize = "CustomNormalize"; -class MS_CORE_API CustomNormalize : public PrimitiveC { +class MIND_API CustomNormalize : public BaseOperator { public: - CustomNormalize() : PrimitiveC(kNameCustomNormalize) {} - ~CustomNormalize() = default; - MS_DECLARE_PARENT(CustomNormalize, PrimitiveC); + MIND_API_BASE_MEMBER(CustomNormalize); + CustomNormalize() : BaseOperator(kNameCustomNormalize) {} void Init() const {} }; -AbstractBasePtr CustomNormalizeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CustomNormalizeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/custom_predict.cc b/mindspore/core/ops/custom_predict.cc index f17d1a59d0..8a3cb3d13c 100644 --- a/mindspore/core/ops/custom_predict.cc +++ b/mindspore/core/ops/custom_predict.cc @@ -18,15 +18,19 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(CustomPredict, PrimitiveC, BaseOperator); void CustomPredict::Init(const int64_t output_num, const float weight_threshold) { this->set_output_num(output_num); this->set_weight_threshold(weight_threshold); } -void CustomPredict::set_output_num(const int64_t output_num) { (void)this->AddAttr(kOutputNum, MakeValue(output_num)); } +void CustomPredict::set_output_num(const int64_t output_num) { + (void)this->AddAttr(kOutputNum, api::MakeValue(output_num)); +} int64_t CustomPredict::get_output_num() const { auto value_ptr = this->GetAttr(kOutputNum); @@ -34,7 +38,7 @@ int64_t CustomPredict::get_output_num() const { } void CustomPredict::set_weight_threshold(const float weight_threshold) { - (void)this->AddAttr(kWeightThreshold, MakeValue(weight_threshold)); + (void)this->AddAttr(kWeightThreshold, api::MakeValue(weight_threshold)); } float CustomPredict::get_weight_threshold() const { diff --git a/mindspore/core/ops/custom_predict.h b/mindspore/core/ops/custom_predict.h index 9fa4026156..17eed1c0a8 100644 --- a/mindspore/core/ops/custom_predict.h +++ b/mindspore/core/ops/custom_predict.h @@ -18,27 +18,24 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCustomPredict = "CustomPredict"; -class MS_CORE_API CustomPredict : public PrimitiveC { +class MIND_API CustomPredict : public BaseOperator { public: - CustomPredict() : PrimitiveC(kNameCustomPredict) {} - ~CustomPredict() = default; - MS_DECLARE_PARENT(CustomPredict, PrimitiveC); + MIND_API_BASE_MEMBER(CustomPredict); + CustomPredict() : BaseOperator(kNameCustomPredict) {} void Init(const int64_t output_num, const float weight_threshold); void set_output_num(const int64_t output_num); void set_weight_threshold(const float weight_threshold); int64_t get_output_num() const; float get_weight_threshold() const; }; -AbstractBasePtr CustomPredictInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CustomPredictInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/data_format_dim_map.cc b/mindspore/core/ops/data_format_dim_map.cc index d6291a5065..ba88663ed3 100644 --- a/mindspore/core/ops/data_format_dim_map.cc +++ b/mindspore/core/ops/data_format_dim_map.cc @@ -22,6 +22,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -49,6 +51,7 @@ TypePtr DataFormatDimMapInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/data_format_dim_map.h b/mindspore/core/ops/data_format_dim_map.h index 6fdfdde4b3..e5981f1e73 100644 --- a/mindspore/core/ops/data_format_dim_map.h +++ b/mindspore/core/ops/data_format_dim_map.h @@ -19,21 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDataFormatDimMap = "DataFormatDimMap"; -class MS_CORE_API DataFormatDimMap : public PrimitiveC { +class MIND_API DataFormatDimMap : public BaseOperator { public: - DataFormatDimMap() : PrimitiveC(kNameDataFormatDimMap) { InitIOName({"x"}, {"output"}); } - ~DataFormatDimMap() = default; - MS_DECLARE_PARENT(DataFormatDimMap, PrimitiveC); + MIND_API_BASE_MEMBER(DataFormatDimMap); + DataFormatDimMap() : BaseOperator(kNameDataFormatDimMap) { InitIOName({"x"}, {"output"}); } }; -AbstractBasePtr DataFormatDimMapInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DataFormatDimMapInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kDataFormatDimMapPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/depend.cc b/mindspore/core/ops/depend.cc index 10dc1d14c4..ec240b5516 100644 --- a/mindspore/core/ops/depend.cc +++ b/mindspore/core/ops/depend.cc @@ -15,9 +15,12 @@ */ #include "ops/depend.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Depend, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameDepend, Depend); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/depend.h b/mindspore/core/ops/depend.h index d9c17e152b..dcedd0c327 100644 --- a/mindspore/core/ops/depend.h +++ b/mindspore/core/ops/depend.h @@ -19,29 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameDepend = "Depend"; /// \brief Depend defined Depend operator prototype. -class MS_CORE_API Depend : public PrimitiveC { +class MIND_API Depend : public BaseOperator { public: + MIND_API_BASE_MEMBER(Depend); /// \brief Constructor. - Depend() : PrimitiveC(kNameDepend) {} - - /// \brief Destructor. - ~Depend() = default; - - MS_DECLARE_PARENT(Depend, PrimitiveC); + Depend() : BaseOperator(kNameDepend) {} /// \brief Method to init the op's attributes. void Init() const {} }; -AbstractBasePtr DependInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DependInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimDepend = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/depth_to_space.cc b/mindspore/core/ops/depth_to_space.cc index 10bf5ef8e5..5c912dc5cb 100644 --- a/mindspore/core/ops/depth_to_space.cc +++ b/mindspore/core/ops/depth_to_space.cc @@ -22,18 +22,19 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void DepthToSpace::set_block_size(const int64_t block_size) { CheckAndConvertUtils::Check(kBlockSize, block_size, kGreaterEqual, 2, this->name()); - (void)this->AddAttr(kBlockSize, MakeValue(block_size)); + (void)this->AddAttr(kBlockSize, api::MakeValue(block_size)); } int64_t DepthToSpace::get_block_size() const { return GetValue(GetAttr(kBlockSize)); } void DepthToSpace::set_format(const Format &format) { int64_t f = format; - (void)this->AddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } Format DepthToSpace::get_format() const { return Format(GetValue(GetAttr(kFormat))); } @@ -112,6 +113,7 @@ TypePtr DepthToSpaceInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = DepthToSpaceInferType(primitive, input_args); diff --git a/mindspore/core/ops/depth_to_space.h b/mindspore/core/ops/depth_to_space.h index e38e911a1c..6143a5af2e 100644 --- a/mindspore/core/ops/depth_to_space.h +++ b/mindspore/core/ops/depth_to_space.h @@ -21,22 +21,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameDepthToSpace = "DepthToSpace"; /// \brief Rearrange blocks of depth data into spatial dimensions. /// Refer to Python API @ref mindspore.ops.DepthToSpace for more details. -class MS_CORE_API DepthToSpace : public PrimitiveC { +class MIND_API DepthToSpace : public BaseOperator { public: + MIND_API_BASE_MEMBER(DepthToSpace); /// \brief Constructor. - DepthToSpace() : PrimitiveC(kNameDepthToSpace) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~DepthToSpace() = default; - MS_DECLARE_PARENT(DepthToSpace, PrimitiveC); + DepthToSpace() : BaseOperator(kNameDepthToSpace) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.DepthToSpace for the inputs. void Init(const int64_t block_size, const Format &format = NCHW); /// \brief Set block_size. @@ -53,8 +51,8 @@ class MS_CORE_API DepthToSpace : public PrimitiveC { Format get_format() const; }; -AbstractBasePtr DepthToSpaceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DepthToSpaceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/detection_post_process.cc b/mindspore/core/ops/detection_post_process.cc index 3ebf4e05a7..6d2e10b0be 100644 --- a/mindspore/core/ops/detection_post_process.cc +++ b/mindspore/core/ops/detection_post_process.cc @@ -17,9 +17,11 @@ #include "ops/detection_post_process.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(DetectionPostProcess, PrimitiveC, BaseOperator); void DetectionPostProcess::Init(const int64_t inputSize, const std::vector &scale, const float NmsIouThreshold, const float NmsScoreThreshold, const int64_t MaxDetections, const int64_t DetectionsPerClass, const int64_t MaxClassesPerDetection, @@ -39,7 +41,7 @@ void DetectionPostProcess::Init(const int64_t inputSize, const std::vectorAddAttr(kInputSize, MakeValue(inputSize)); + (void)this->AddAttr(kInputSize, api::MakeValue(inputSize)); } int64_t DetectionPostProcess::get_input_size() const { @@ -47,14 +49,16 @@ int64_t DetectionPostProcess::get_input_size() const { return GetValue(value_ptr); } -void DetectionPostProcess::set_scale(const std::vector &scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } +void DetectionPostProcess::set_scale(const std::vector &scale) { + (void)this->AddAttr(kScale, api::MakeValue(scale)); +} std::vector DetectionPostProcess::get_scale() const { auto value_ptr = this->GetAttr(kScale); return GetValue>(value_ptr); } void DetectionPostProcess::set_nms_iou_threshold(const float NmsIouThreshold) { - (void)this->AddAttr(kNmsIouThreshold, MakeValue(NmsIouThreshold)); + (void)this->AddAttr(kNmsIouThreshold, api::MakeValue(NmsIouThreshold)); } float DetectionPostProcess::get_nms_iou_threshold() const { auto value_ptr = this->GetAttr(kNmsIouThreshold); @@ -62,7 +66,7 @@ float DetectionPostProcess::get_nms_iou_threshold() const { } void DetectionPostProcess::set_nms_score_threshold(const float NmsScoreThreshold) { - (void)this->AddAttr(kNmsScoreThreshold, MakeValue(NmsScoreThreshold)); + (void)this->AddAttr(kNmsScoreThreshold, api::MakeValue(NmsScoreThreshold)); } float DetectionPostProcess::get_nms_score_threshold() const { auto value_ptr = this->GetAttr(kNmsScoreThreshold); @@ -70,12 +74,12 @@ float DetectionPostProcess::get_nms_score_threshold() const { } void DetectionPostProcess::set_max_detections(const int64_t MaxDetections) { - (void)this->AddAttr(kMaxDetections, MakeValue(MaxDetections)); + (void)this->AddAttr(kMaxDetections, api::MakeValue(MaxDetections)); } int64_t DetectionPostProcess::get_max_detections() const { return GetValue(GetAttr(kMaxDetections)); } void DetectionPostProcess::set_detections_per_class(const int64_t DetectionsPerClass) { - (void)this->AddAttr(kDetectionsPerClass, MakeValue(DetectionsPerClass)); + (void)this->AddAttr(kDetectionsPerClass, api::MakeValue(DetectionsPerClass)); } int64_t DetectionPostProcess::get_detections_per_class() const { auto value_ptr = this->GetAttr(kDetectionsPerClass); @@ -83,18 +87,18 @@ int64_t DetectionPostProcess::get_detections_per_class() const { } void DetectionPostProcess::set_max_classes_per_detection(const int64_t MaxClassesPerDetection) { - (void)this->AddAttr(kMaxClassesPerDetection, MakeValue(MaxClassesPerDetection)); + (void)this->AddAttr(kMaxClassesPerDetection, api::MakeValue(MaxClassesPerDetection)); } int64_t DetectionPostProcess::get_max_classes_per_detection() const { return GetValue(GetAttr(kMaxClassesPerDetection)); } void DetectionPostProcess::set_num_classes(const int64_t NumClasses) { - (void)this->AddAttr(kNumClasses, MakeValue(NumClasses)); + (void)this->AddAttr(kNumClasses, api::MakeValue(NumClasses)); } int64_t DetectionPostProcess::get_num_classes() const { return GetValue(GetAttr(kNumClasses)); } void DetectionPostProcess::set_use_regular_nms(const bool UseRegularNms) { - (void)this->AddAttr(kUseRegularNms, MakeValue(UseRegularNms)); + (void)this->AddAttr(kUseRegularNms, api::MakeValue(UseRegularNms)); } bool DetectionPostProcess::get_use_regular_nms() const { auto value_ptr = this->GetAttr(kUseRegularNms); @@ -102,7 +106,7 @@ bool DetectionPostProcess::get_use_regular_nms() const { } void DetectionPostProcess::set_out_quantized(const bool OutQuantized) { - (void)this->AddAttr(kOutQuantized, MakeValue(OutQuantized)); + (void)this->AddAttr(kOutQuantized, api::MakeValue(OutQuantized)); } bool DetectionPostProcess::get_out_quantized() const { auto value_ptr = this->GetAttr(kOutQuantized); @@ -110,7 +114,7 @@ bool DetectionPostProcess::get_out_quantized() const { } void DetectionPostProcess::set_format(const Format &format) { int64_t f = format; - (void)this->AddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } Format DetectionPostProcess::get_format() const { return Format(GetValue(GetAttr(kFormat))); } diff --git a/mindspore/core/ops/detection_post_process.h b/mindspore/core/ops/detection_post_process.h index d776ce9a22..5de492b420 100644 --- a/mindspore/core/ops/detection_post_process.h +++ b/mindspore/core/ops/detection_post_process.h @@ -20,18 +20,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameDetectionPostProcess = "DetectionPostProcess"; -class MS_CORE_API DetectionPostProcess : public PrimitiveC { +class MIND_API DetectionPostProcess : public BaseOperator { public: - DetectionPostProcess() : PrimitiveC(kNameDetectionPostProcess) {} - ~DetectionPostProcess() = default; - MS_DECLARE_PARENT(DetectionPostProcess, PrimitiveC); + MIND_API_BASE_MEMBER(DetectionPostProcess); + DetectionPostProcess() : BaseOperator(kNameDetectionPostProcess) {} void Init(const int64_t inputSize, const std::vector &scale, const float NmsIouThreshold, const float NmsScoreThreshold, const int64_t MaxDetections, const int64_t DetectionsPerClass, const int64_t MaxClassesPerDetection, const int64_t NumClasses, const bool UseRegularNms, @@ -62,8 +61,8 @@ class MS_CORE_API DetectionPostProcess : public PrimitiveC { bool get_out_quantized() const; Format get_format() const; }; -AbstractBasePtr DetectionPostProcessInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DetectionPostProcessInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/diag.cc b/mindspore/core/ops/diag.cc index 5c3ddaccfc..b20f960d06 100644 --- a/mindspore/core/ops/diag.cc +++ b/mindspore/core/ops/diag.cc @@ -20,6 +20,8 @@ #include #include "ops/diag.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +42,8 @@ TypePtr PartInferType(const PrimitivePtr &primitive, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(Diag, PrimitiveC, BaseOperator); AbstractBasePtr DiagInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/diag.h b/mindspore/core/ops/diag.h index 5ae1f152fc..21276ef9b8 100644 --- a/mindspore/core/ops/diag.h +++ b/mindspore/core/ops/diag.h @@ -18,24 +18,21 @@ #define MINDSPORE_CORE_OPS_DIAG_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Constructs a diagonal tensor with a given diagonal values. /// Refer to Python API @ref mindspore.ops.Diag for more details. -class MS_CORE_API Diag : public PrimitiveC { +class MIND_API Diag : public BaseOperator { public: + MIND_API_BASE_MEMBER(Diag); /// \brief Constructor. - Diag() : PrimitiveC(prim::kPrimDiag->name()) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~Diag() = default; - MS_DECLARE_PARENT(Diag, PrimitiveC); + Diag() : BaseOperator("Diag") { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr DiagInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DiagInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/diag_part.cc b/mindspore/core/ops/diag_part.cc index f4b03d8815..b53be6c52e 100644 --- a/mindspore/core/ops/diag_part.cc +++ b/mindspore/core/ops/diag_part.cc @@ -20,6 +20,8 @@ #include #include "ops/diag_part.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -49,6 +51,8 @@ TypePtr DiagPartInferType(const PrimitivePtr &primitive, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(DiagPart, PrimitiveC, BaseOperator); AbstractBasePtr DiagPartInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/diag_part.h b/mindspore/core/ops/diag_part.h index b140ba878f..aa1e9a41e3 100644 --- a/mindspore/core/ops/diag_part.h +++ b/mindspore/core/ops/diag_part.h @@ -18,24 +18,21 @@ #define MINDSPORE_CORE_OPS_DIAG_PART_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Extracts the diagonal part from given tensor. /// Refer to Python API @ref mindspore.ops.DiagPart for more details. -class MS_CORE_API DiagPart : public PrimitiveC { +class MIND_API DiagPart : public BaseOperator { public: + MIND_API_BASE_MEMBER(DiagPart); /// \brief Constructor. - DiagPart() : PrimitiveC(prim::kPrimDiagPart->name()) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~DiagPart() = default; - MS_DECLARE_PARENT(DiagPart, PrimitiveC); + DiagPart() : BaseOperator("DiagPart") { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr DiagPartInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DiagPartInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/div.cc b/mindspore/core/ops/div.cc index 8cb5aa93c1..d58131cc12 100644 --- a/mindspore/core/ops/div.cc +++ b/mindspore/core/ops/div.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Div, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameDiv, Div); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/div.h b/mindspore/core/ops/div.h index d8cab26b11..e85f0e49b9 100644 --- a/mindspore/core/ops/div.h +++ b/mindspore/core/ops/div.h @@ -19,28 +19,25 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDiv = "Div"; /// \brief Computes the quotient of dividing the first input tensor by the second input tensor element-wise. /// Refer to Python API @ref mindspore.ops.Div for more details. -class MS_CORE_API Div : public PrimitiveC { +class MIND_API Div : public BaseOperator { public: + MIND_API_BASE_MEMBER(Div); /// \brief Constructor. - Div() : PrimitiveC(kNameDiv) { InitIOName({"x", "y"}, {"output"}); } - explicit Div(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Div() = default; - MS_DECLARE_PARENT(Div, PrimitiveC); + Div() : BaseOperator(kNameDiv) { InitIOName({"x", "y"}, {"output"}); } + explicit Div(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Div for the inputs. void Init() const {} }; -AbstractBasePtr DivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/div_no_nan.cc b/mindspore/core/ops/div_no_nan.cc index 48c4f690c3..a13ac8f0c7 100644 --- a/mindspore/core/ops/div_no_nan.cc +++ b/mindspore/core/ops/div_no_nan.cc @@ -23,6 +23,9 @@ #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -160,6 +163,7 @@ ValuePtr DivNoNanInferValue(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_shape = DivNoNanInferShape(primitive, input_args); diff --git a/mindspore/core/ops/div_no_nan.h b/mindspore/core/ops/div_no_nan.h index c6b12ddf09..046ccb2bff 100644 --- a/mindspore/core/ops/div_no_nan.h +++ b/mindspore/core/ops/div_no_nan.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDivNoNan = "DivNoNan"; -class DivNoNan : public PrimitiveC { +class DivNoNan : public BaseOperator { public: - DivNoNan() : PrimitiveC(kNameDivNoNan) { InitIOName({"x1", "x2"}, {"y"}); } - ~DivNoNan() = default; - MS_DECLARE_PARENT(DivNoNan, PrimitiveC); + MIND_API_BASE_MEMBER(DivNoNan); + DivNoNan() : BaseOperator(kNameDivNoNan) { InitIOName({"x1", "x2"}, {"y"}); } }; -AbstractBasePtr DivNoNanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DivNoNanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimDivNoNanPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dropout.cc b/mindspore/core/ops/dropout.cc index 9621d32846..6c35045e67 100644 --- a/mindspore/core/ops/dropout.cc +++ b/mindspore/core/ops/dropout.cc @@ -21,14 +21,16 @@ #include "ops/dropout.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Dropout, PrimitiveC, BaseOperator); void Dropout::Init(const float keep_prob) { this->set_keep_prob(keep_prob); } void Dropout::set_keep_prob(const float keep_prob) { CheckAndConvertUtils::CheckInRange(kKeepProb, keep_prob, kIncludeRight, {0.0, 1.0}, this->name()); - (void)this->AddAttr(kKeepProb, MakeValue(keep_prob)); + (void)this->AddAttr(kKeepProb, api::MakeValue(keep_prob)); } float Dropout::get_keep_prob() const { diff --git a/mindspore/core/ops/dropout.h b/mindspore/core/ops/dropout.h index 88550242d1..c5a8c140ec 100644 --- a/mindspore/core/ops/dropout.h +++ b/mindspore/core/ops/dropout.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDropout = "Dropout"; /// \brief During training, randomly zeroes some of the elements of the input tensor with probability 1-keep_prob //// from a Bernoulli distribution. Refer to Python API @ref mindspore.ops.Dropout for more details. -class MS_CORE_API Dropout : public PrimitiveC { +class MIND_API Dropout : public BaseOperator { public: + MIND_API_BASE_MEMBER(Dropout); /// \brief Constructor. - Dropout() : PrimitiveC(kNameDropout) {} - /// \brief Destructor. - ~Dropout() = default; - MS_DECLARE_PARENT(Dropout, PrimitiveC); + Dropout() : BaseOperator(kNameDropout) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Dropout for the inputs. void Init(const float keep_prob = 0.5); /// \brief Set keep_prob. @@ -44,8 +41,8 @@ class MS_CORE_API Dropout : public PrimitiveC { /// \return keep_prob. float get_keep_prob() const; }; -AbstractBasePtr DropoutInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DropoutInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_DROPOUT_H_ diff --git a/mindspore/core/ops/dropout_do_mask.cc b/mindspore/core/ops/dropout_do_mask.cc index cf9db440c6..81c1ea5dea 100644 --- a/mindspore/core/ops/dropout_do_mask.cc +++ b/mindspore/core/ops/dropout_do_mask.cc @@ -24,6 +24,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -120,6 +121,7 @@ TypePtr DropoutDoMaskInferType(const PrimitivePtr &primitive, const std::vector< } } // namespace +MIND_API_BASE_IMPL(DropoutDoMask, PrimitiveC, BaseOperator); AbstractBasePtr DropoutDoMaskInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/dropout_do_mask.h b/mindspore/core/ops/dropout_do_mask.h index b29d6ef819..ea812ec7a1 100644 --- a/mindspore/core/ops/dropout_do_mask.h +++ b/mindspore/core/ops/dropout_do_mask.h @@ -19,26 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Applies dropout mask on the input tensor. /// Refer to Python API @ref mindspore.ops.DropoutDoMask for more details. -class MS_CORE_API DropoutDoMask : public PrimitiveC { +class MIND_API DropoutDoMask : public BaseOperator { public: + MIND_API_BASE_MEMBER(DropoutDoMask); /// \brief Constructor. - DropoutDoMask() : PrimitiveC(prim::kPrimDropoutDoMask->name()) {} - /// \brief Destructor. - ~DropoutDoMask() = default; - MS_DECLARE_PARENT(DropoutDoMask, PrimitiveC); + DropoutDoMask() : BaseOperator("DropoutDoMask") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.DropoutDoMask for the inputs. void Init() const {} }; -AbstractBasePtr DropoutDoMaskInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DropoutDoMaskInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dropout_gen_mask.cc b/mindspore/core/ops/dropout_gen_mask.cc index 823dfe83c4..8afc8046ba 100644 --- a/mindspore/core/ops/dropout_gen_mask.cc +++ b/mindspore/core/ops/dropout_gen_mask.cc @@ -20,11 +20,13 @@ #include #include #include +#include #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -164,6 +166,7 @@ TypePtr DropoutGenMaskInferType(const PrimitivePtr &primitive, const std::vector } } // namespace +MIND_API_BASE_IMPL(DropoutGenMask, PrimitiveC, BaseOperator); AbstractBasePtr DropoutGenMaskInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/dropout_gen_mask.h b/mindspore/core/ops/dropout_gen_mask.h index d681a41d8d..4eeabe8fa0 100644 --- a/mindspore/core/ops/dropout_gen_mask.h +++ b/mindspore/core/ops/dropout_gen_mask.h @@ -19,26 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Generates the mask value for the input shape. /// Refer to Python API @ref mindspore.ops.DropoutGenMask for more details. -class MS_CORE_API DropoutGenMask : public PrimitiveC { +class MIND_API DropoutGenMask : public BaseOperator { public: + MIND_API_BASE_MEMBER(DropoutGenMask); /// \brief Constructor. - DropoutGenMask() : PrimitiveC(prim::kPrimDropoutGenMask->name()) {} - /// \brief Destructor. - ~DropoutGenMask() = default; - MS_DECLARE_PARENT(DropoutGenMask, PrimitiveC); + DropoutGenMask() : BaseOperator("DropoutGenMask") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.DropoutGenMask for the inputs. void Init() const {} }; -AbstractBasePtr DropoutGenMaskInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DropoutGenMaskInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dtype.cc b/mindspore/core/ops/dtype.cc index e9914b1c49..eb22e44551 100644 --- a/mindspore/core/ops/dtype.cc +++ b/mindspore/core/ops/dtype.cc @@ -23,10 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/abstract_value.h" -#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(DType, PrimitiveC, BaseOperator); ValuePtr DTypeInferValue(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); auto op_name = primitive->name(); diff --git a/mindspore/core/ops/dtype.h b/mindspore/core/ops/dtype.h index d8cff677df..50e73b1784 100644 --- a/mindspore/core/ops/dtype.h +++ b/mindspore/core/ops/dtype.h @@ -20,21 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns the data type of the input tensor as mindspore.dtype. /// Refer to Python API @ref mindspore.ops.DType for more details. -class MS_CORE_API DType : public PrimitiveC { +class MIND_API DType : public BaseOperator { public: + MIND_API_BASE_MEMBER(DType); /// \brief Constructor. - DType() : PrimitiveC(prim::kPrimDType->name()) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~DType() = default; - MS_DECLARE_PARENT(DType, PrimitiveC); + DType() : BaseOperator("DType") { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.DType for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/dynamic_broadcast_gradient_args.cc b/mindspore/core/ops/dynamic_broadcast_gradient_args.cc index 6a37d59a76..0225456725 100644 --- a/mindspore/core/ops/dynamic_broadcast_gradient_args.cc +++ b/mindspore/core/ops/dynamic_broadcast_gradient_args.cc @@ -23,6 +23,8 @@ #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -69,6 +71,7 @@ abstract::TupleShapePtr Infer(const PrimitivePtr &primitive, const std::vector &input_args) { auto types = std::vector{kInt64, kInt64}; diff --git a/mindspore/core/ops/dynamic_broadcast_gradient_args.h b/mindspore/core/ops/dynamic_broadcast_gradient_args.h index 7645197825..1172be7707 100644 --- a/mindspore/core/ops/dynamic_broadcast_gradient_args.h +++ b/mindspore/core/ops/dynamic_broadcast_gradient_args.h @@ -18,20 +18,21 @@ #define MINDSPORE_CORE_OPS_DYNAMIC_BROADCAST_GRADIENT_ARGS_H_ #include #include -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API DynamicBroadcastGradientArgs : public PrimitiveC { +constexpr auto kNameDynamicBroadcastGradientArgs = "DynamicBroadcastGradientArgs"; +class MIND_API DynamicBroadcastGradientArgs : public BaseOperator { public: - DynamicBroadcastGradientArgs() : PrimitiveC(prim::kPrimDynamicBroadcastGradientArgs->name()) {} - ~DynamicBroadcastGradientArgs() = default; - MS_DECLARE_PARENT(DynamicBroadcastGradientArgs, PrimitiveC); + MIND_API_BASE_MEMBER(DynamicBroadcastGradientArgs); + DynamicBroadcastGradientArgs() : BaseOperator(kNameDynamicBroadcastGradientArgs) {} void Init() const {} }; -AbstractBasePtr DynamicBroadcastGradientArgsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DynamicBroadcastGradientArgsInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dynamic_broadcast_to.cc b/mindspore/core/ops/dynamic_broadcast_to.cc index b449a25e83..27379a8cf5 100644 --- a/mindspore/core/ops/dynamic_broadcast_to.cc +++ b/mindspore/core/ops/dynamic_broadcast_to.cc @@ -18,6 +18,7 @@ #include #include "utils/check_convert_utils.h" +#include "ops/op_utils.h" namespace mindspore { namespace ops { @@ -84,6 +85,7 @@ TypePtr DynamicBroadcastToInferType(const PrimitivePtr &prim, const std::vector< } } // namespace +MIND_API_BASE_IMPL(DynamicBroadcastTo, PrimitiveC, BaseOperator); AbstractBasePtr DynamicBroadcastToInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(DynamicBroadcastToInferShape(primitive, input_args), diff --git a/mindspore/core/ops/dynamic_broadcast_to.h b/mindspore/core/ops/dynamic_broadcast_to.h index eee6d3ca74..42462f4e0a 100644 --- a/mindspore/core/ops/dynamic_broadcast_to.h +++ b/mindspore/core/ops/dynamic_broadcast_to.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -class DynamicBroadcastTo : public PrimitiveC { +class DynamicBroadcastTo : public BaseOperator { public: - DynamicBroadcastTo() : PrimitiveC(prim::kPrimDynamicBroadcastTo->name()) { InitIOName({"x", "shape"}, {"y"}); } - ~DynamicBroadcastTo() = default; - MS_DECLARE_PARENT(DynamicBroadcastTo, PrimitiveC); + MIND_API_BASE_MEMBER(DynamicBroadcastTo); + DynamicBroadcastTo() : BaseOperator("DynamicBroadcastTo") { InitIOName({"x", "shape"}, {"y"}); } void Init() {} }; -AbstractBasePtr DynamicBroadcastToInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DynamicBroadcastToInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimDynamicBroadcastToPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dynamic_quant.cc b/mindspore/core/ops/dynamic_quant.cc index 4907dca9e9..6d7f1a5644 100644 --- a/mindspore/core/ops/dynamic_quant.cc +++ b/mindspore/core/ops/dynamic_quant.cc @@ -15,15 +15,20 @@ */ #include "ops/dynamic_quant.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void DynamicQuant::set_symmetric(const bool symmetric) { (void)AddAttr(kSymmetric, MakeValue(symmetric)); } +MIND_API_BASE_IMPL(DynamicQuant, PrimitiveC, BaseOperator); +void DynamicQuant::set_symmetric(const bool symmetric) { (void)AddAttr(kSymmetric, api::MakeValue(symmetric)); } bool DynamicQuant::get_symmetric() const { auto value_ptr = this->GetAttr(kSymmetric); return GetValue(value_ptr); } -void DynamicQuant::set_dst_type(const int64_t dst_type) { (void)AddAttr(kDstType, MakeValue(dst_type)); } +void DynamicQuant::set_dst_type(const int64_t dst_type) { (void)AddAttr(kDstType, api::MakeValue(dst_type)); } int64_t DynamicQuant::get_dst_type() const { return GetValue(GetAttr(kDstType)); } void DynamicQuant::Init(const bool symmetric, const int64_t dst_type) { this->set_symmetric(symmetric); diff --git a/mindspore/core/ops/dynamic_quant.h b/mindspore/core/ops/dynamic_quant.h index 1b7b254b1a..ade36b4f92 100644 --- a/mindspore/core/ops/dynamic_quant.h +++ b/mindspore/core/ops/dynamic_quant.h @@ -22,25 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDynamicQuant = "DynamicQuant"; /// \brief the DynamicQuant operator prototype. -class MS_CORE_API DynamicQuant : public PrimitiveC { +class MIND_API DynamicQuant : public BaseOperator { public: + MIND_API_BASE_MEMBER(DynamicQuant); /// \brief Constructor. - DynamicQuant() : PrimitiveC(kNameDynamicQuant) {} - - /// \brief Destructor. - ~DynamicQuant() = default; - - MS_DECLARE_PARENT(DynamicQuant, PrimitiveC); + DynamicQuant() : BaseOperator(kNameDynamicQuant) {} /// \brief Method to init the op's attributes. /// @@ -68,8 +62,8 @@ class MS_CORE_API DynamicQuant : public PrimitiveC { /// \return the data type of output. int64_t get_dst_type() const; }; -AbstractBasePtr DynamicQuantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DynamicQuantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dynamic_resize_nearest_neighbor.cc b/mindspore/core/ops/dynamic_resize_nearest_neighbor.cc index 66d42b6aaa..54d143607f 100644 --- a/mindspore/core/ops/dynamic_resize_nearest_neighbor.cc +++ b/mindspore/core/ops/dynamic_resize_nearest_neighbor.cc @@ -23,6 +23,7 @@ #include "ops/dynamic_resize_nearest_neighbor.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -118,6 +119,8 @@ TypePtr DynamicResizeNearestNeighborInferType(const PrimitivePtr &prim, return CheckAndConvertUtils::CheckTensorTypeValid("x", input_args[0]->BuildType(), valid_types, prim->name()); } } // namespace + +MIND_API_BASE_IMPL(DynamicResizeNearestNeighbor, PrimitiveC, BaseOperator); AbstractBasePtr DynamicResizeNearestNeighborInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto prim_name = primitive->name(); diff --git a/mindspore/core/ops/dynamic_resize_nearest_neighbor.h b/mindspore/core/ops/dynamic_resize_nearest_neighbor.h index fb0fb46983..932c7ceb11 100644 --- a/mindspore/core/ops/dynamic_resize_nearest_neighbor.h +++ b/mindspore/core/ops/dynamic_resize_nearest_neighbor.h @@ -19,19 +19,17 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDynamicResizeNearestNeighbor = "DynamicResizeNearestNeighbor"; -class DynamicResizeNearestNeighbor : public PrimitiveC { +class DynamicResizeNearestNeighbor : public BaseOperator { public: - DynamicResizeNearestNeighbor() : PrimitiveC(kNameDynamicResizeNearestNeighbor) {} - ~DynamicResizeNearestNeighbor() = default; - MS_DECLARE_PARENT(DynamicResizeNearestNeighbor, PrimitiveC); + MIND_API_BASE_MEMBER(DynamicResizeNearestNeighbor); + DynamicResizeNearestNeighbor() : BaseOperator(kNameDynamicResizeNearestNeighbor) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/dynamic_shape.h b/mindspore/core/ops/dynamic_shape.h index cf0b075b85..cdd38f855e 100644 --- a/mindspore/core/ops/dynamic_shape.h +++ b/mindspore/core/ops/dynamic_shape.h @@ -15,15 +15,15 @@ */ #ifndef MINDSPORE_CORE_OPS_DYNAMIC_SHAPE_H_ #define MINDSPORE_CORE_OPS_DYNAMIC_SHAPE_H_ -#include "ops/primitive_c.h" -#include "base/core_ops.h" +#include "ops/base_operator.h" + namespace mindspore { namespace ops { -class DynamicShape : public PrimitiveC { +constexpr auto kNameDynamicShape = "DynamicShape"; +class MIND_API DynamicShape : public BaseOperator { public: - DynamicShape() : PrimitiveC(prim::kPrimDynamicShape->name()) {} - ~DynamicShape() = default; - MS_DECLARE_PARENT(DynamicShape, PrimitiveC); + MIND_API_BASE_MEMBER(DynamicShape); + DynamicShape() : BaseOperator(kNameDynamicShape) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/einsum.cc b/mindspore/core/ops/einsum.cc index 57bdc03808..83b16be182 100644 --- a/mindspore/core/ops/einsum.cc +++ b/mindspore/core/ops/einsum.cc @@ -24,6 +24,8 @@ #include "ir/dtype/tensor_type.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -217,9 +219,10 @@ static void element_map_shape(const std::string &prim_name, const std::vectorset_equation(equation); } -void Einsum::set_equation(const std::string &equation) { (void)this->AddAttr(kEquation, MakeValue(equation)); } +void Einsum::set_equation(const std::string &equation) { (void)this->AddAttr(kEquation, api::MakeValue(equation)); } std::string Einsum::get_equation() const { auto value_ptr = this->GetAttr(kEquation); diff --git a/mindspore/core/ops/einsum.h b/mindspore/core/ops/einsum.h index 538ce3f0ec..c32d2a4055 100644 --- a/mindspore/core/ops/einsum.h +++ b/mindspore/core/ops/einsum.h @@ -20,9 +20,9 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { @@ -30,13 +30,11 @@ constexpr auto kNameEinsum = "Einsum"; /// \brief . /// Refer to Python API @ref mindspore.ops.Einsum for more details. -class Einsum : public PrimitiveC { +class MIND_API Einsum : public BaseOperator { public: + MIND_API_BASE_MEMBER(Einsum); /// \brief Constructor. - Einsum() : PrimitiveC(kNameEinsum) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Einsum() = default; - MS_DECLARE_PARENT(Einsum, PrimitiveC); + Einsum() : BaseOperator(kNameEinsum) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Einsum for the inputs. void Init(const std::string &equation); /// \brief Set equation. @@ -47,8 +45,8 @@ class Einsum : public PrimitiveC { std::string get_equation() const; }; -AbstractBasePtr EinsumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr EinsumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/eltwise.cc b/mindspore/core/ops/eltwise.cc index c9e9421a75..2c598bf886 100644 --- a/mindspore/core/ops/eltwise.cc +++ b/mindspore/core/ops/eltwise.cc @@ -17,13 +17,15 @@ #include "ops/eltwise.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Eltwise, PrimitiveC, BaseOperator); void Eltwise::Init(const EltwiseMode &mode) { this->set_mode(mode); } void Eltwise::set_mode(const EltwiseMode &mode) { int64_t m = mode; - (void)this->AddAttr(kMode, MakeValue(m)); + (void)this->AddAttr(kMode, api::MakeValue(m)); } EltwiseMode Eltwise::get_mode() const { auto value_ptr = this->GetAttr(kMode); diff --git a/mindspore/core/ops/eltwise.h b/mindspore/core/ops/eltwise.h index 8c02461539..c525b3f0f1 100644 --- a/mindspore/core/ops/eltwise.h +++ b/mindspore/core/ops/eltwise.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_ELTWISE_H_ #define MINDSPORE_CORE_OPS_ELTWISE_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameEltwise = "Eltwise"; /// \brief Eltwise defined Element-wise operator prototype. -class MS_CORE_API Eltwise : public PrimitiveC { +class MIND_API Eltwise : public BaseOperator { public: + MIND_API_BASE_MEMBER(Eltwise); /// \brief Constructor. - Eltwise() : PrimitiveC(kNameEltwise) {} - - /// \brief Destructor. - ~Eltwise() = default; - - MS_DECLARE_PARENT(Eltwise, PrimitiveC); + Eltwise() : BaseOperator(kNameEltwise) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/elu.cc b/mindspore/core/ops/elu.cc index 3e3ab44037..4df85f63bf 100644 --- a/mindspore/core/ops/elu.cc +++ b/mindspore/core/ops/elu.cc @@ -25,6 +25,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -54,10 +55,12 @@ TypePtr EluInferType(const PrimitivePtr &prim, const std::vectorset_alpha(alpha); } void Elu::set_alpha(const float alpha) { - (void)AddAttr(kAlpha, MakeValue(CheckAndConvertUtils::CheckValue(kAlpha, alpha, kEqual, 1.0, name()))); + (void)AddAttr(kAlpha, api::MakeValue(CheckAndConvertUtils::CheckValue(kAlpha, alpha, kEqual, 1.0, name()))); } float Elu::get_alpha() const { diff --git a/mindspore/core/ops/elu.h b/mindspore/core/ops/elu.h index bff1d4f562..c85c4413de 100644 --- a/mindspore/core/ops/elu.h +++ b/mindspore/core/ops/elu.h @@ -18,22 +18,18 @@ #define MINDSPORE_CORE_OPS_ELU_H_ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameElu = "Elu"; /// \brief Calculate exponential linearity. Refer to Python API @ref mindspore.ops.Elu for more details. -class MS_CORE_API Elu : public PrimitiveC { +class MIND_API Elu : public BaseOperator { public: + MIND_API_BASE_MEMBER(Elu); /// \brief Constructor. - Elu() : PrimitiveC(kNameElu) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Elu() = default; - MS_DECLARE_PARENT(Elu, PrimitiveC); + Elu() : BaseOperator(kNameElu) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Elu for the inputs. void Init(const float alpha = 0.0); /// \brief Set alpha. @@ -43,8 +39,8 @@ class MS_CORE_API Elu : public PrimitiveC { /// \return alpha. float get_alpha() const; }; -AbstractBasePtr EluInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr EluInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimElu = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/embedding_lookup.cc b/mindspore/core/ops/embedding_lookup.cc index 2fea8ff726..543ebc2271 100644 --- a/mindspore/core/ops/embedding_lookup.cc +++ b/mindspore/core/ops/embedding_lookup.cc @@ -19,13 +19,15 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(EmbeddingLookup, PrimitiveC, BaseOperator); void EmbeddingLookup::Init(const bool setattr_flag) { this->set_setattr_flag(setattr_flag); } void EmbeddingLookup::set_setattr_flag(const bool setattr_flag) { - (void)this->AddAttr(kSetattrFlag, MakeValue(setattr_flag)); + (void)this->AddAttr(kSetattrFlag, api::MakeValue(setattr_flag)); } bool EmbeddingLookup::get_setattr_flag() const { diff --git a/mindspore/core/ops/embedding_lookup.h b/mindspore/core/ops/embedding_lookup.h index 90d01d7bb0..4cd99d940d 100644 --- a/mindspore/core/ops/embedding_lookup.h +++ b/mindspore/core/ops/embedding_lookup.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameEmbeddingLookup = "EmbeddingLookup"; /// \brief Returns a slice of input tensor based on the specified indices. /// Refer to Python API @ref mindspore.ops.EmbeddingLookup for more details. -class MS_CORE_API EmbeddingLookup : public PrimitiveC { +class MIND_API EmbeddingLookup : public BaseOperator { public: + MIND_API_BASE_MEMBER(EmbeddingLookup); /// \brief Constructor. - EmbeddingLookup() : PrimitiveC(kNameEmbeddingLookup) { InitIOName({"params", "indices", "offset"}, {"output"}); } - /// \brief Destructor. - ~EmbeddingLookup() = default; - MS_DECLARE_PARENT(EmbeddingLookup, PrimitiveC); + EmbeddingLookup() : BaseOperator(kNameEmbeddingLookup) { InitIOName({"params", "indices", "offset"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.EmbeddingLookup for the inputs. void Init(const bool setattr_flag = true); /// \brief Set setattr_flag. @@ -45,8 +42,8 @@ class MS_CORE_API EmbeddingLookup : public PrimitiveC { /// \return setattr_flag. bool get_setattr_flag() const; }; -AbstractBasePtr EmbeddingLookupInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr EmbeddingLookupInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/equal.cc b/mindspore/core/ops/equal.cc index 2d1d273a82..8dd340dc8b 100644 --- a/mindspore/core/ops/equal.cc +++ b/mindspore/core/ops/equal.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Equal, PrimitiveC, BaseOperator); AbstractBasePtr EqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/equal.h b/mindspore/core/ops/equal.h index d6dff56d6d..02222b7cc8 100644 --- a/mindspore/core/ops/equal.h +++ b/mindspore/core/ops/equal.h @@ -19,28 +19,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameEqual = prim::kEqual; +constexpr auto kNameEqual = "Equal"; /// \brief Computes the equivalence between two tensors element-wise. /// Refer to Python API @ref mindspore.ops.Equal for more details. -class MS_CORE_API Equal : public PrimitiveC { +class MIND_API Equal : public BaseOperator { public: + MIND_API_BASE_MEMBER(Equal); /// \brief Constructor. - Equal() : PrimitiveC(prim::kPrimEqual->name()) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Equal() = default; - MS_DECLARE_PARENT(Equal, PrimitiveC); + Equal() : BaseOperator("Equal") { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Equal for the inputs. void Init() const {} }; -AbstractBasePtr EqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr EqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/erf.cc b/mindspore/core/ops/erf.cc index 32a77cbb58..4c6f533873 100644 --- a/mindspore/core/ops/erf.cc +++ b/mindspore/core/ops/erf.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -57,6 +58,8 @@ TypePtr ErfInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/erf.h b/mindspore/core/ops/erf.h index cc824d71b0..1ea67f277a 100644 --- a/mindspore/core/ops/erf.h +++ b/mindspore/core/ops/erf.h @@ -18,22 +18,21 @@ #define MINDSPORE_CORE_OPS_ERF_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameErf = "Erf"; -class MS_CORE_API Erf : public PrimitiveC { +class MIND_API Erf : public BaseOperator { public: - Erf() : PrimitiveC(kNameErf) { InitIOName({"x"}, {"y"}); } - ~Erf() = default; - MS_DECLARE_PARENT(Erf, PrimitiveC); + MIND_API_BASE_MEMBER(Erf); + /// \brief Constructor. + Erf() : BaseOperator(kNameErf) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr ErfInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ErfInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimErf = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/erfc.cc b/mindspore/core/ops/erfc.cc index cacdf27255..5faa5114ce 100644 --- a/mindspore/core/ops/erfc.cc +++ b/mindspore/core/ops/erfc.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -57,6 +58,8 @@ TypePtr ErfcInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/erfc.h b/mindspore/core/ops/erfc.h index 5666a01c15..781146ca22 100644 --- a/mindspore/core/ops/erfc.h +++ b/mindspore/core/ops/erfc.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_ERFC_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameErfc = "Erfc"; -class Erfc : public PrimitiveC { +class Erfc : public BaseOperator { public: - Erfc() : PrimitiveC(kNameErfc) { InitIOName({"x"}, {"y"}); } - ~Erfc() = default; - MS_DECLARE_PARENT(Erfc, PrimitiveC); + MIND_API_BASE_MEMBER(Erfc); + Erfc() : BaseOperator(kNameErfc) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr ErfcInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ErfcInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimErfc = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/erfinv.cc b/mindspore/core/ops/erfinv.cc index d1d301e552..a123617412 100644 --- a/mindspore/core/ops/erfinv.cc +++ b/mindspore/core/ops/erfinv.cc @@ -20,6 +20,8 @@ #include #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -37,6 +39,8 @@ TypePtr ErfinvInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/erfinv.h b/mindspore/core/ops/erfinv.h index 7c59d1a49b..54c40d0258 100644 --- a/mindspore/core/ops/erfinv.h +++ b/mindspore/core/ops/erfinv.h @@ -18,25 +18,22 @@ #define MINDSPORE_CORE_OPS_ERFINV_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameErfinv = "Erfinv"; /// \brief Computes the inverse error function of input. Refer to Python API @ref mindspore.ops.Erfinv for more details. -class Erfinv : public PrimitiveC { +class MIND_API Erfinv : public BaseOperator { public: + MIND_API_BASE_MEMBER(Erfinv); /// \brief Constructor. - Erfinv() : PrimitiveC(kNameErfinv) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~Erfinv() = default; - MS_DECLARE_PARENT(Erfinv, PrimitiveC); + Erfinv() : BaseOperator(kNameErfinv) { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr ErfinvInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ErfinvInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/exp.cc b/mindspore/core/ops/exp.cc index c4a172f74f..f5e10c2e74 100644 --- a/mindspore/core/ops/exp.cc +++ b/mindspore/core/ops/exp.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -54,6 +55,8 @@ TypePtr ExpInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(Exp, PrimitiveC, BaseOperator); AbstractBasePtr ExpInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(ExpInferShape(primitive, input_args), ExpInferType(primitive, input_args)); diff --git a/mindspore/core/ops/exp.h b/mindspore/core/ops/exp.h index 0d5b501078..80f46415e2 100644 --- a/mindspore/core/ops/exp.h +++ b/mindspore/core/ops/exp.h @@ -20,28 +20,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameExp = prim::kExp; +constexpr auto kNameExp = "Exp"; /// \brief Returns exponential of a tensor element-wise. Refer to Python API @ref mindspore.ops.Exp for more details. -class MS_CORE_API Exp : public PrimitiveC { +class MIND_API Exp : public BaseOperator { public: + MIND_API_BASE_MEMBER(Exp); /// \brief Constructor. - Exp() : PrimitiveC(prim::kPrimExp->name()) { InitIOName({"x"}, {"y"}); } - explicit Exp(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Exp() = default; - MS_DECLARE_PARENT(Exp, PrimitiveC); + Exp() : BaseOperator("Exp") { InitIOName({"x"}, {"y"}); } + explicit Exp(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Exp for the inputs. void Init() const {} }; -AbstractBasePtr ExpInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ExpInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/expand_dims.cc b/mindspore/core/ops/expand_dims.cc index 275d95c2ba..ccb39024ee 100644 --- a/mindspore/core/ops/expand_dims.cc +++ b/mindspore/core/ops/expand_dims.cc @@ -24,9 +24,11 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/log_adapter.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ExpandDims, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameExpandDims, ExpandDims); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/expand_dims.h b/mindspore/core/ops/expand_dims.h index 3f5fed1541..1f1cc62e18 100644 --- a/mindspore/core/ops/expand_dims.h +++ b/mindspore/core/ops/expand_dims.h @@ -18,27 +18,24 @@ #define MINDSPORE_CORE_OPS_EXPAND_DIMS_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameExpandDims = "ExpandDims"; /// \brief Adds an additional dimension to ‘input_x` at the given axis. /// Refer to Python API @ref mindspore.ops.ExpandDims for more details. -class MS_CORE_API ExpandDims : public PrimitiveC { +class MIND_API ExpandDims : public BaseOperator { public: + MIND_API_BASE_MEMBER(ExpandDims); /// \brief Constructor. - ExpandDims() : PrimitiveC(kNameExpandDims) { InitIOName({"x", "axis"}, {"output"}); } - /// \brief Destructor. - ~ExpandDims() = default; - MS_DECLARE_PARENT(ExpandDims, PrimitiveC); + ExpandDims() : BaseOperator(kNameExpandDims) { InitIOName({"x", "axis"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ExpandDims for the inputs. void Init() const {} }; -AbstractBasePtr ExpandDimsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ExpandDimsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimExpandDims = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/expm1.cc b/mindspore/core/ops/expm1.cc index 4e57042583..df0dd543f6 100644 --- a/mindspore/core/ops/expm1.cc +++ b/mindspore/core/ops/expm1.cc @@ -19,6 +19,9 @@ #include #include #include "ops/expm1.h" +#include "utils/check_convert_utils.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +45,8 @@ TypePtr Expm1InferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/expm1.h b/mindspore/core/ops/expm1.h index 1c2fe33711..58cfa5207b 100644 --- a/mindspore/core/ops/expm1.h +++ b/mindspore/core/ops/expm1.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_EXPM1_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameExpm1 = "Expm1"; -class Expm1 : public PrimitiveC { +class MIND_API Expm1 : public BaseOperator { public: - Expm1() : PrimitiveC(kNameExpm1) { InitIOName({"x"}, {"output"}); } - ~Expm1() = default; - MS_DECLARE_PARENT(Expm1, PrimitiveC); + MIND_API_BASE_MEMBER(Expm1); + Expm1() : BaseOperator(kNameExpm1) { InitIOName({"x"}, {"output"}); } void Init() {} }; -AbstractBasePtr Expm1Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr Expm1Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimExpm1Ptr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/extract_volume_patches.cc b/mindspore/core/ops/extract_volume_patches.cc index c77e0da9af..59f14e8279 100644 --- a/mindspore/core/ops/extract_volume_patches.cc +++ b/mindspore/core/ops/extract_volume_patches.cc @@ -16,6 +16,8 @@ #include "ops/extract_volume_patches.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -87,6 +89,8 @@ TypePtr ExtractVolumePatchesInferType(const PrimitivePtr &prim, const std::vecto return CheckAndConvertUtils::CheckTensorTypeValid("x", input_args[0]->BuildType(), valid_types, prim->name()); } } // namespace + +MIND_API_BASE_IMPL(ExtractVolumePatches, PrimitiveC, BaseOperator); AbstractBasePtr ExtractVolumePatchesInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/extract_volume_patches.h b/mindspore/core/ops/extract_volume_patches.h index 0e65e2c473..07fcf26b21 100644 --- a/mindspore/core/ops/extract_volume_patches.h +++ b/mindspore/core/ops/extract_volume_patches.h @@ -21,26 +21,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameExtractVolumePatches = "ExtractVolumePatches"; /// \brief Extract patches from input and put them in the "depth" output dimension. /// Refer to Python API @ref mindspore.ops.ExtractVolumePatches for more details. -class MS_CORE_API ExtractVolumePatches : public PrimitiveC { +class MIND_API ExtractVolumePatches : public BaseOperator { public: + MIND_API_BASE_MEMBER(ExtractVolumePatches); /// \brief Constructor. - ExtractVolumePatches() : PrimitiveC(kNameExtractVolumePatches) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~ExtractVolumePatches() = default; - MS_DECLARE_PARENT(ExtractVolumePatches, PrimitiveC); + ExtractVolumePatches() : BaseOperator(kNameExtractVolumePatches) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr ExtractVolumePatchesInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ExtractVolumePatchesInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimExtractVolumePatchesPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fake_quant_with_min_max_vars.cc b/mindspore/core/ops/fake_quant_with_min_max_vars.cc index 93b68fce43..871bf44806 100644 --- a/mindspore/core/ops/fake_quant_with_min_max_vars.cc +++ b/mindspore/core/ops/fake_quant_with_min_max_vars.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -31,7 +32,7 @@ void FakeQuantWithMinMaxVars::Init(const bool narrow_range, const int64_t num_bi } void FakeQuantWithMinMaxVars::set_narrow_range(const bool narrow_range) { - (void)this->AddAttr(kNarrowRange, MakeValue(narrow_range)); + (void)this->AddAttr(kNarrowRange, api::MakeValue(narrow_range)); } bool FakeQuantWithMinMaxVars::get_narrow_range() const { @@ -40,7 +41,7 @@ bool FakeQuantWithMinMaxVars::get_narrow_range() const { } void FakeQuantWithMinMaxVars::set_num_bits(const int64_t num_bits) { - (void)this->AddAttr(kNumBits, MakeValue(num_bits)); + (void)this->AddAttr(kNumBits, api::MakeValue(num_bits)); } int64_t FakeQuantWithMinMaxVars::get_num_bits() const { @@ -48,6 +49,7 @@ int64_t FakeQuantWithMinMaxVars::get_num_bits() const { return GetValue(value_ptr); } +MIND_API_BASE_IMPL(FakeQuantWithMinMaxVars, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameFakeQuantWithMinMaxVars, FakeQuantWithMinMaxVars); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fake_quant_with_min_max_vars.h b/mindspore/core/ops/fake_quant_with_min_max_vars.h index 97de0ea237..0446edc788 100644 --- a/mindspore/core/ops/fake_quant_with_min_max_vars.h +++ b/mindspore/core/ops/fake_quant_with_min_max_vars.h @@ -18,22 +18,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFakeQuantWithMinMaxVars = "FakeQuantWithMinMaxVars"; /// \brief Fake-quantize the input by minimum and maximum. /// Refer to Python API @ref mindspore.ops.FakeQuantWithMinMaxVars for more details. -class MS_CORE_API FakeQuantWithMinMaxVars : public PrimitiveC { +class MIND_API FakeQuantWithMinMaxVars : public BaseOperator { public: + MIND_API_BASE_MEMBER(FakeQuantWithMinMaxVars); /// \brief Constructor. - FakeQuantWithMinMaxVars() : PrimitiveC(kNameFakeQuantWithMinMaxVars) {} - /// \brief Destructor. - ~FakeQuantWithMinMaxVars() = default; - MS_DECLARE_PARENT(FakeQuantWithMinMaxVars, PrimitiveC); + FakeQuantWithMinMaxVars() : BaseOperator(kNameFakeQuantWithMinMaxVars) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.FakeQuantWithMinMaxVars for the inputs. void Init(const bool narrow_range = false, const int64_t num_bits = 8); /// \brief Set narrow_range. @@ -49,8 +46,9 @@ class MS_CORE_API FakeQuantWithMinMaxVars : public PrimitiveC { /// \return num_bits. int64_t get_num_bits() const; }; -AbstractBasePtr FakeQuantWithMinMaxVarsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FakeQuantWithMinMaxVarsInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.cc b/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.cc index 231299fc90..213f9affa1 100644 --- a/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.cc +++ b/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.cc @@ -16,19 +16,22 @@ #include "ops/fake_quant_with_min_max_vars_per_channel.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(FakeQuantWithMinMaxVarsPerChannel, PrimitiveC, BaseOperator); void FakeQuantWithMinMaxVarsPerChannel::Init(const int64_t num_bits, const bool narrow_range) { this->set_num_bits(num_bits); this->set_narrow_range(narrow_range); } void FakeQuantWithMinMaxVarsPerChannel::set_num_bits(const int64_t num_bits) { (void)CheckAndConvertUtils::CheckInteger(kNumBits, num_bits, kGreaterThan, 0, this->name()); - (void)this->AddAttr(kNumBits, MakeValue(num_bits)); + (void)this->AddAttr(kNumBits, api::MakeValue(num_bits)); } void FakeQuantWithMinMaxVarsPerChannel::set_narrow_range(const bool narrow_range) { - (void)this->AddAttr(kNarrowRange, MakeValue(narrow_range)); + (void)this->AddAttr(kNarrowRange, api::MakeValue(narrow_range)); } int64_t FakeQuantWithMinMaxVarsPerChannel::get_num_bits() const { auto value_ptr = GetAttr(kNumBits); diff --git a/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.h b/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.h index 33396a2d00..ba77c499ca 100644 --- a/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.h +++ b/mindspore/core/ops/fake_quant_with_min_max_vars_per_channel.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFakeQuantWithMinMaxVarsPerChannel = "FakeQuantWithMinMaxVarsPerChannel"; /// \brief Fake-quantize the input and one of shape: [d], [b, d], [b, h, w, d] by per-channel minimum and maximum. /// Refer to Python API @ref mindspore.ops.FakeQuantWithMinMaxVarsPerChannel for more details. -class MS_CORE_API FakeQuantWithMinMaxVarsPerChannel : public PrimitiveC { +class MIND_API FakeQuantWithMinMaxVarsPerChannel : public BaseOperator { public: + MIND_API_BASE_MEMBER(FakeQuantWithMinMaxVarsPerChannel); /// \brief Constructor. - FakeQuantWithMinMaxVarsPerChannel() : PrimitiveC(kNameFakeQuantWithMinMaxVarsPerChannel) {} - /// \brief Destructor. - ~FakeQuantWithMinMaxVarsPerChannel() = default; - MS_DECLARE_PARENT(FakeQuantWithMinMaxVarsPerChannel, PrimitiveC); + FakeQuantWithMinMaxVarsPerChannel() : BaseOperator(kNameFakeQuantWithMinMaxVarsPerChannel) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.FakeQuantWithMinMaxVarsPerChannel /// for the inputs. void Init(const int64_t num_bits = 8, const bool narrow_range = false); @@ -53,9 +50,9 @@ class MS_CORE_API FakeQuantWithMinMaxVarsPerChannel : public PrimitiveC { bool get_narrow_range() const; }; -AbstractBasePtr FakeQuantWithMinMaxVarsPerChannelInfer(const abstract::AnalysisEnginePtr &, - const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FakeQuantWithMinMaxVarsPerChannelInfer( + const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fast_gelu.cc b/mindspore/core/ops/fast_gelu.cc index 8d8946436b..5c772e2332 100644 --- a/mindspore/core/ops/fast_gelu.cc +++ b/mindspore/core/ops/fast_gelu.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -52,6 +53,7 @@ TypePtr FastGeLUInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = FastGeLUInferType(primitive, input_args); diff --git a/mindspore/core/ops/fast_gelu.h b/mindspore/core/ops/fast_gelu.h index e33137b980..3719e783f2 100644 --- a/mindspore/core/ops/fast_gelu.h +++ b/mindspore/core/ops/fast_gelu.h @@ -22,24 +22,20 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFastGeLU = "FastGeLU"; -class FastGeLU : public PrimitiveC { +class FastGeLU : public BaseOperator { public: - FastGeLU() : PrimitiveC(kNameFastGeLU) { InitIOName({"x"}, {"y"}); } - - ~FastGeLU() = default; - - MS_DECLARE_PARENT(FastGeLU, PrimitiveC); + MIND_API_BASE_MEMBER(FastGeLU); + FastGeLU() : BaseOperator(kNameFastGeLU) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr FastGeLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FastGeLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimFastGeLUPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/fft_imag.cc b/mindspore/core/ops/fft_imag.cc index e8d8154827..779af47e9e 100644 --- a/mindspore/core/ops/fft_imag.cc +++ b/mindspore/core/ops/fft_imag.cc @@ -18,9 +18,11 @@ #include #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(FftImag, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameFftImag, FftImag); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fft_imag.h b/mindspore/core/ops/fft_imag.h index 8a542dd77c..e4a0ff0765 100644 --- a/mindspore/core/ops/fft_imag.h +++ b/mindspore/core/ops/fft_imag.h @@ -20,29 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFftImag = "FftImag"; /// \brief FftImag defined Imaginary-part acquisition operator prototype. -class MS_CORE_API FftImag : public PrimitiveC { +class MIND_API FftImag : public BaseOperator { public: + MIND_API_BASE_MEMBER(FftImag); /// \brief Constructor. - FftImag() : PrimitiveC(kNameFftImag) {} - - /// \brief Destructor. - ~FftImag() = default; - - MS_DECLARE_PARENT(FftImag, PrimitiveC); + FftImag() : BaseOperator(kNameFftImag) {} /// \brief Method to init the op's attributes. void Init() const {} }; -AbstractBasePtr FftImagInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FftImagInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fft_real.cc b/mindspore/core/ops/fft_real.cc index fcfd76eb95..5d3df6191a 100644 --- a/mindspore/core/ops/fft_real.cc +++ b/mindspore/core/ops/fft_real.cc @@ -21,9 +21,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(FftReal, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameFftReal, FftReal); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fft_real.h b/mindspore/core/ops/fft_real.h index 82c748296d..568a3f1476 100644 --- a/mindspore/core/ops/fft_real.h +++ b/mindspore/core/ops/fft_real.h @@ -19,29 +19,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFftReal = "FftReal"; /// \brief FftReal defined Real-part acquisition operator prototype. -class MS_CORE_API FftReal : public PrimitiveC { +class MIND_API FftReal : public BaseOperator { public: + MIND_API_BASE_MEMBER(FftReal); /// \brief Constructor. - FftReal() : PrimitiveC(kNameFftReal) {} - - /// \brief Destructor. - ~FftReal() = default; - MS_DECLARE_PARENT(FftReal, PrimitiveC); + FftReal() : BaseOperator(kNameFftReal) {} /// \brief Method to init the op's attributes. void Init() const {} }; -AbstractBasePtr FftRealInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FftRealInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fill.cc b/mindspore/core/ops/fill.cc index 680804f226..2ee87d9038 100644 --- a/mindspore/core/ops/fill.cc +++ b/mindspore/core/ops/fill.cc @@ -19,9 +19,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Fill, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameFill, Fill); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fill.h b/mindspore/core/ops/fill.h index 74a4d2dd41..b2fc5641c0 100644 --- a/mindspore/core/ops/fill.h +++ b/mindspore/core/ops/fill.h @@ -19,27 +19,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFill = "Fill"; /// \brief Creates a tensor filled with a scalar value. Refer to Python API @ref mindspore.ops.Fill for more details. -class MS_CORE_API Fill : public PrimitiveC { +class MIND_API Fill : public BaseOperator { public: + MIND_API_BASE_MEMBER(Fill); /// \brief Constructor. - Fill() : PrimitiveC(kNameFill) {} - /// \brief Destructor. - ~Fill() = default; - MS_DECLARE_PARENT(Fill, PrimitiveC); + Fill() : BaseOperator(kNameFill) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Fill for the inputs. void Init() const {} }; -AbstractBasePtr FillInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FillInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fill_v2.cc b/mindspore/core/ops/fill_v2.cc index 20a54fd027..09c2d12fe5 100644 --- a/mindspore/core/ops/fill_v2.cc +++ b/mindspore/core/ops/fill_v2.cc @@ -25,9 +25,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(FillV2, PrimitiveC, BaseOperator); abstract::ShapePtr FillV2InferShape(const PrimitivePtr &primitive, const std::vector &input_args) { if (!input_args[0]->isa()) { MS_EXCEPTION(TypeError) << "Input[0] only support tensor!"; diff --git a/mindspore/core/ops/fill_v2.h b/mindspore/core/ops/fill_v2.h index 6e7714945d..d414bfddac 100644 --- a/mindspore/core/ops/fill_v2.h +++ b/mindspore/core/ops/fill_v2.h @@ -21,23 +21,20 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFillV2 = "FillV2"; -class MS_CORE_API FillV2 : public PrimitiveC { +class MIND_API FillV2 : public BaseOperator { public: - FillV2() : PrimitiveC(kNameFillV2) { InitIOName({"shape", "value"}, {"y"}); } - ~FillV2() = default; - MS_DECLARE_PARENT(FillV2, PrimitiveC); + MIND_API_BASE_MEMBER(FillV2); + FillV2() : BaseOperator(kNameFillV2) { InitIOName({"shape", "value"}, {"y"}); } }; -AbstractBasePtr FillV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FillV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimFillV2Ptr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/flatten.cc b/mindspore/core/ops/flatten.cc index 42633e58e3..38b42a0303 100644 --- a/mindspore/core/ops/flatten.cc +++ b/mindspore/core/ops/flatten.cc @@ -17,6 +17,8 @@ #include #include "ops/flatten.h" #include "utils/check_convert_utils.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -69,6 +71,7 @@ TypePtr FlattenInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = FlattenInferShape(primitive, input_args); diff --git a/mindspore/core/ops/flatten.h b/mindspore/core/ops/flatten.h index 5009211ed6..922f2d070b 100644 --- a/mindspore/core/ops/flatten.h +++ b/mindspore/core/ops/flatten.h @@ -19,27 +19,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFlatten = "Flatten"; /// \brief Flattens a tensor without changing its batch size on the 0-th axis. /// Refer to Python API @ref mindspore.ops.Flatten for more details. -class MS_CORE_API Flatten : public PrimitiveC { +class MIND_API Flatten : public BaseOperator { public: + MIND_API_BASE_MEMBER(Flatten); /// \brief Constructor. - Flatten() : PrimitiveC(kNameFlatten) {} - /// \brief Destructor. - ~Flatten() = default; - MS_DECLARE_PARENT(Flatten, PrimitiveC); + Flatten() : BaseOperator(kNameFlatten) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Flatten for the inputs. void Init() const {} }; -AbstractBasePtr FlattenInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FlattenInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/floor.cc b/mindspore/core/ops/floor.cc index 7734d23a54..2ba66c8184 100644 --- a/mindspore/core/ops/floor.cc +++ b/mindspore/core/ops/floor.cc @@ -22,6 +22,7 @@ #include #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,8 @@ TypePtr FloorInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/floor.h b/mindspore/core/ops/floor.h index a1a31de71a..52e4ef59da 100644 --- a/mindspore/core/ops/floor.h +++ b/mindspore/core/ops/floor.h @@ -21,27 +21,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFloor = "Floor"; /// \brief Rounds a tensor down to the closest integer element-wise. /// Refer to Python API @ref mindspore.ops.Floor for more details. -class MS_CORE_API Floor : public PrimitiveC { +class MIND_API Floor : public BaseOperator { public: + MIND_API_BASE_MEMBER(Floor); /// \brief Constructor. - Floor() : PrimitiveC(kNameFloor) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Floor() = default; - MS_DECLARE_PARENT(Floor, PrimitiveC); + Floor() : BaseOperator(kNameFloor) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Floor for the inputs. void Init() const {} }; -AbstractBasePtr FloorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FloorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimFloorPtr = std::shared_ptr; diff --git a/mindspore/core/ops/floor_div.cc b/mindspore/core/ops/floor_div.cc index 8e21b5948a..29ecb6ae48 100644 --- a/mindspore/core/ops/floor_div.cc +++ b/mindspore/core/ops/floor_div.cc @@ -25,6 +25,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -67,6 +68,7 @@ TypePtr FloorDivInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/floor_div.h b/mindspore/core/ops/floor_div.h index cd100ef0fb..6361aa38a3 100644 --- a/mindspore/core/ops/floor_div.h +++ b/mindspore/core/ops/floor_div.h @@ -18,31 +18,28 @@ #define MINDSPORE_CORE_OPS_FLOOR_DIV_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFloorDiv = "FloorDiv"; /// \brief Divides the first input tensor by the second input tensor element-wise and round down to the closest integer. /// Refer to Python API @ref mindspore.ops.FloorDiv for more details. -class MS_CORE_API FloorDiv : public PrimitiveC { +class MIND_API FloorDiv : public BaseOperator { public: + MIND_API_BASE_MEMBER(FloorDiv); /// \brief Constructor. - FloorDiv() : PrimitiveC(kNameFloorDiv) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~FloorDiv() = default; - MS_DECLARE_PARENT(FloorDiv, PrimitiveC); + FloorDiv() : BaseOperator(kNameFloorDiv) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.FloorDiv for the inputs. void Init() const {} - AbstractBasePtr FloorDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); + abstract::AbstractBasePtr FloorDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimFloorDivPtr = std::shared_ptr; }; -AbstractBasePtr FloorDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FloorDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/floor_mod.cc b/mindspore/core/ops/floor_mod.cc index 17cb876e0e..f55df6adef 100644 --- a/mindspore/core/ops/floor_mod.cc +++ b/mindspore/core/ops/floor_mod.cc @@ -25,11 +25,13 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" -#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { namespace { +using mindspore::Complex; + abstract::ShapePtr FloorModInferShape(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); auto prim_name = primitive->name(); @@ -89,6 +91,7 @@ TypePtr FloorModInferType(const PrimitivePtr &primitive, const std::vector &input_args) { auto type = FloorModInferType(primitive, input_args); diff --git a/mindspore/core/ops/floor_mod.h b/mindspore/core/ops/floor_mod.h index 17de41f66f..eefbc8c90b 100644 --- a/mindspore/core/ops/floor_mod.h +++ b/mindspore/core/ops/floor_mod.h @@ -17,27 +17,24 @@ #ifndef MINDSPORE_CORE_OPS_FLOOR_MOD_H_ #define MINDSPORE_CORE_OPS_FLOOR_MOD_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFloorMod = "FloorMod"; /// \brief Computes the remainder of division element-wise. /// Refer to Python API @ref mindspore.ops.FloorMod for more details. -class MS_CORE_API FloorMod : public PrimitiveC { +class MIND_API FloorMod : public BaseOperator { public: + MIND_API_BASE_MEMBER(FloorMod); /// \brief Constructor. - FloorMod() : PrimitiveC(kNameFloorMod) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~FloorMod() = default; - MS_DECLARE_PARENT(FloorMod, PrimitiveC); + FloorMod() : BaseOperator(kNameFloorMod) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.FloorMod for the inputs. void Init() const {} }; -AbstractBasePtr FloorModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FloorModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fused_batch_norm.cc b/mindspore/core/ops/fused_batch_norm.cc index 68f197a17e..34ecb55f37 100644 --- a/mindspore/core/ops/fused_batch_norm.cc +++ b/mindspore/core/ops/fused_batch_norm.cc @@ -18,20 +18,22 @@ #include "ops/fused_batch_norm.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(FusedBatchNorm, PrimitiveC, BaseOperator); void FusedBatchNorm::Init(const int64_t mode, const float epsilon, const float momentum) { this->set_mode(mode); this->set_epsilon(epsilon); this->set_momentum(momentum); } -void FusedBatchNorm::set_mode(const int64_t mode) { (void)this->AddAttr(kMode, MakeValue(mode)); } +void FusedBatchNorm::set_mode(const int64_t mode) { (void)this->AddAttr(kMode, api::MakeValue(mode)); } -void FusedBatchNorm::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, MakeValue(epsilon)); } +void FusedBatchNorm::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, api::MakeValue(epsilon)); } -void FusedBatchNorm::set_momentum(const float momentum) { (void)this->AddAttr(kMomentum, MakeValue(momentum)); } +void FusedBatchNorm::set_momentum(const float momentum) { (void)this->AddAttr(kMomentum, api::MakeValue(momentum)); } int64_t FusedBatchNorm::get_mode() const { auto value_ptr = this->GetAttr(kMode); diff --git a/mindspore/core/ops/fused_batch_norm.h b/mindspore/core/ops/fused_batch_norm.h index ac83458101..4f350a1c35 100644 --- a/mindspore/core/ops/fused_batch_norm.h +++ b/mindspore/core/ops/fused_batch_norm.h @@ -20,27 +20,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFusedBatchNorm = "FusedBatchNorm"; /// \brief FusedBatchNorm defined Enhanced BatchNorm operator prototype. -class MS_CORE_API FusedBatchNorm : public PrimitiveC { +class MIND_API FusedBatchNorm : public BaseOperator { public: + MIND_API_BASE_MEMBER(FusedBatchNorm); /// \brief Constructor. - FusedBatchNorm() : PrimitiveC(kNameFusedBatchNorm) { + FusedBatchNorm() : BaseOperator(kNameFusedBatchNorm) { InitIOName({"x", "scale", "b", "mean", "variance"}, {"y", "running_mean", "running_variance", "save_mean", "save_inv_variance"}); } - /// \brief Destructor. - ~FusedBatchNorm() = default; - - MS_DECLARE_PARENT(FusedBatchNorm, PrimitiveC); - /// \brief Method to init the op's attributes. /// /// \param[in] mode Define the mode of batchnorm, which is useless. diff --git a/mindspore/core/ops/fusion/activation.cc b/mindspore/core/ops/fusion/activation.cc index 8436de53b0..9db81d8097 100644 --- a/mindspore/core/ops/fusion/activation.cc +++ b/mindspore/core/ops/fusion/activation.cc @@ -20,18 +20,20 @@ #include #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Activation::set_alpha(const float alpha) { (void)this->AddAttr(kAlpha, MakeValue(alpha)); } +MIND_API_BASE_IMPL(Activation, PrimitiveC, BaseOperator); +void Activation::set_alpha(const float alpha) { (void)this->AddAttr(kAlpha, api::MakeValue(alpha)); } -void Activation::set_min_val(const float min_val) { (void)this->AddAttr(kMinVal, MakeValue(min_val)); } +void Activation::set_min_val(const float min_val) { (void)this->AddAttr(kMinVal, api::MakeValue(min_val)); } -void Activation::set_max_val(const float max_val) { (void)this->AddAttr(kMaxVal, MakeValue(max_val)); } +void Activation::set_max_val(const float max_val) { (void)this->AddAttr(kMaxVal, api::MakeValue(max_val)); } void Activation::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } float Activation::get_alpha() const { @@ -54,7 +56,7 @@ ActivationType Activation::get_activation_type() const { return ActivationType(GetValue(value_ptr)); } -void Activation::set_approximate(bool approximate) { (void)this->AddAttr(kApproximate, MakeValue(approximate)); } +void Activation::set_approximate(bool approximate) { (void)this->AddAttr(kApproximate, api::MakeValue(approximate)); } bool Activation::get_approximate() const { auto value_ptr = this->GetAttr(kApproximate); diff --git a/mindspore/core/ops/fusion/activation.h b/mindspore/core/ops/fusion/activation.h index b4ce98b25c..860bede450 100644 --- a/mindspore/core/ops/fusion/activation.h +++ b/mindspore/core/ops/fusion/activation.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_ACTIVATION_H_ #define MINDSPORE_CORE_OPS_ACTIVATION_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameActivation = "Activation"; /// \brief Activation defined Activation operator prototype of lite. -class MS_CORE_API Activation : public PrimitiveC { +class MIND_API Activation : public BaseOperator { public: + MIND_API_BASE_MEMBER(Activation); /// \brief Constructor. - Activation() : PrimitiveC(kNameActivation) {} - - /// \brief Destructor. - ~Activation() = default; - - MS_DECLARE_PARENT(Activation, PrimitiveC); + Activation() : BaseOperator(kNameActivation) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/fusion/add_fusion.cc b/mindspore/core/ops/fusion/add_fusion.cc index d28343f9ac..b1e8746db0 100644 --- a/mindspore/core/ops/fusion/add_fusion.cc +++ b/mindspore/core/ops/fusion/add_fusion.cc @@ -22,12 +22,14 @@ #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void AddFusion::set_activation_type(const ActivationType activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType AddFusion::get_activation_type() const { auto value_ptr = GetAttr(kActivationType); @@ -35,6 +37,7 @@ ActivationType AddFusion::get_activation_type() const { } void AddFusion::Init(const ActivationType activation_type) { this->set_activation_type(activation_type); } +MIND_API_BASE_IMPL(AddFusion, PrimitiveC, Add); REGISTER_PRIMITIVE_C(kNameAddFusion, AddFusion); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/add_fusion.h b/mindspore/core/ops/fusion/add_fusion.h index 2c5ff37480..c7e295adb4 100644 --- a/mindspore/core/ops/fusion/add_fusion.h +++ b/mindspore/core/ops/fusion/add_fusion.h @@ -20,23 +20,18 @@ #include #include "ops/add.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAddFusion = "AddFusion"; /// \brief AddFusion defined Add operator prototype of lite. -class MS_CORE_API AddFusion : public Add { +class MIND_API AddFusion : public Add { public: + MIND_API_BASE_MEMBER(AddFusion); /// \brief Constructor. AddFusion() : Add(kNameAddFusion) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~AddFusion() = default; - - MS_DECLARE_PARENT(AddFusion, Add); - /// \brief Method to init the op's attributes. /// /// \param[in] activation_type Define the activation type. @@ -53,8 +48,8 @@ class MS_CORE_API AddFusion : public Add { ActivationType get_activation_type() const; }; -AbstractBasePtr AddFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AddFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/adder_fusion.cc b/mindspore/core/ops/fusion/adder_fusion.cc index 2df0fb0aaa..7757144b54 100644 --- a/mindspore/core/ops/fusion/adder_fusion.cc +++ b/mindspore/core/ops/fusion/adder_fusion.cc @@ -16,9 +16,11 @@ #include "ops/fusion/adder_fusion.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(AdderFusion, PrimitiveC, Adder); void AdderFusion::Init(const int64_t in_channel, const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &stride, const std::vector &pad_list, const std::vector &dilation, const int64_t group, @@ -37,7 +39,7 @@ void AdderFusion::Init(const int64_t in_channel, const int64_t out_channel, cons void AdderFusion::set_activation_type(const ActivationType activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType AdderFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/adder_fusion.h b/mindspore/core/ops/fusion/adder_fusion.h index 04513ad0bc..e6a17af809 100644 --- a/mindspore/core/ops/fusion/adder_fusion.h +++ b/mindspore/core/ops/fusion/adder_fusion.h @@ -22,23 +22,18 @@ #include #include #include "ops/adder.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAdderFusion = "AdderFusion"; /// \brief AdderFusion defined Adder operator prototype of lite. -class MS_CORE_API AdderFusion : public Adder { +class MIND_API AdderFusion : public Adder { public: + MIND_API_BASE_MEMBER(AdderFusion); /// \brief Constructor. AdderFusion() : Adder(kNameAdderFusion) {} - /// \brief Destructor. - ~AdderFusion() = default; - - MS_DECLARE_PARENT(AdderFusion, Adder); - /// \brief Method to init the op's attributes. /// /// \param[in] in_channel Define the number of input channel. diff --git a/mindspore/core/ops/fusion/arg_max_fusion.cc b/mindspore/core/ops/fusion/arg_max_fusion.cc index d9e4a9498a..874c77d4ed 100644 --- a/mindspore/core/ops/fusion/arg_max_fusion.cc +++ b/mindspore/core/ops/fusion/arg_max_fusion.cc @@ -15,9 +15,12 @@ */ #include "ops/fusion/arg_max_fusion.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ArgMaxFusion, PrimitiveC, ArgMax); void ArgMaxFusion::Init(const bool keep_dims, const bool out_max_value, const int64_t top_k, const int64_t axis) { set_axis(axis); set_keep_dims(keep_dims); @@ -25,11 +28,11 @@ void ArgMaxFusion::Init(const bool keep_dims, const bool out_max_value, const in set_top_k(top_k); } -void ArgMaxFusion::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, MakeValue(keep_dims)); } +void ArgMaxFusion::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, api::MakeValue(keep_dims)); } void ArgMaxFusion::set_out_max_value(const bool out_max_value) { - (void)this->AddAttr(kOutMaxValue, MakeValue(out_max_value)); + (void)this->AddAttr(kOutMaxValue, api::MakeValue(out_max_value)); } -void ArgMaxFusion::set_top_k(const int64_t top_k) { (void)this->AddAttr(kTopK, MakeValue(top_k)); } +void ArgMaxFusion::set_top_k(const int64_t top_k) { (void)this->AddAttr(kTopK, api::MakeValue(top_k)); } bool ArgMaxFusion::get_keep_dims() const { auto keep_dims = GetAttr(kKeepDims); diff --git a/mindspore/core/ops/fusion/arg_max_fusion.h b/mindspore/core/ops/fusion/arg_max_fusion.h index 3f8f93b763..c3ae22e8f4 100644 --- a/mindspore/core/ops/fusion/arg_max_fusion.h +++ b/mindspore/core/ops/fusion/arg_max_fusion.h @@ -25,16 +25,12 @@ namespace mindspore { namespace ops { constexpr auto kNameArgMaxFusion = "ArgMaxFusion"; /// \brief ArgMaxFusion defined ArgMax operator prototype of lite. -class MS_CORE_API ArgMaxFusion : public ArgMax { +class MIND_API ArgMaxFusion : public ArgMax { public: + MIND_API_BASE_MEMBER(ArgMaxFusion); /// \brief Constructor. ArgMaxFusion() : ArgMax(kNameArgMaxFusion) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~ArgMaxFusion() = default; - - MS_DECLARE_PARENT(ArgMaxFusion, ArgMax); - /// \brief Method to init the op's attributes. /// /// \param[in] keep_dims Define a boolean value to indicate the dimension of output is equal to that of input or not. @@ -73,8 +69,8 @@ class MS_CORE_API ArgMaxFusion : public ArgMax { /// \return the number of maximum value along with axis. int64_t get_top_k() const; }; -AbstractBasePtr ArgMaxFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ArgMaxFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimArgMaxFusion = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/arg_min_fusion.cc b/mindspore/core/ops/fusion/arg_min_fusion.cc index 390cda70b8..63ba2fe11d 100644 --- a/mindspore/core/ops/fusion/arg_min_fusion.cc +++ b/mindspore/core/ops/fusion/arg_min_fusion.cc @@ -15,9 +15,12 @@ */ #include "ops/fusion/arg_min_fusion.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ArgMinFusion, PrimitiveC, ArgMin); void ArgMinFusion::Init(bool keep_dims, bool out_max_value, int64_t top_k, int64_t axis) { set_axis(axis); set_keep_dims(keep_dims); @@ -25,9 +28,9 @@ void ArgMinFusion::Init(bool keep_dims, bool out_max_value, int64_t top_k, int64 set_top_k(top_k); } -void ArgMinFusion::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, MakeValue(keep_dims)); } -void ArgMinFusion::set_out_max_value(bool out_max_value) { (void)AddAttr(kOutMaxValue, MakeValue(out_max_value)); } -void ArgMinFusion::set_top_k(int64_t top_k) { (void)this->AddAttr(kTopK, MakeValue(top_k)); } +void ArgMinFusion::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, api::MakeValue(keep_dims)); } +void ArgMinFusion::set_out_max_value(bool out_max_value) { (void)AddAttr(kOutMaxValue, api::MakeValue(out_max_value)); } +void ArgMinFusion::set_top_k(int64_t top_k) { (void)this->AddAttr(kTopK, api::MakeValue(top_k)); } bool ArgMinFusion::get_keep_dims() const { auto keep_dims = GetAttr(kKeepDims); diff --git a/mindspore/core/ops/fusion/arg_min_fusion.h b/mindspore/core/ops/fusion/arg_min_fusion.h index de7e38748b..4a809fe9ba 100644 --- a/mindspore/core/ops/fusion/arg_min_fusion.h +++ b/mindspore/core/ops/fusion/arg_min_fusion.h @@ -25,16 +25,12 @@ namespace mindspore { namespace ops { constexpr auto kNameArgMinFusion = "ArgMinFusion"; /// \brief ArgMinFusion defined ArgMin operator prototype of lite. -class MS_CORE_API ArgMinFusion : public ArgMin { +class MIND_API ArgMinFusion : public ArgMin { public: + MIND_API_BASE_MEMBER(ArgMinFusion); /// \brief Constructor. ArgMinFusion() : ArgMin(kNameArgMinFusion) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~ArgMinFusion() = default; - - MS_DECLARE_PARENT(ArgMinFusion, ArgMin); - /// \brief Method to init the op's attributes. /// /// \param[in] keep_dims Define a boolean value to indicate the dimension of output is equal to that of input or not. @@ -73,8 +69,8 @@ class MS_CORE_API ArgMinFusion : public ArgMin { /// \return the number of minimum value along with axis. int64_t get_top_k() const; }; -AbstractBasePtr ArgMinFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ArgMinFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimArgMinFusion = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/avg_pool_fusion.cc b/mindspore/core/ops/fusion/avg_pool_fusion.cc index db4d69bcdd..87b54ce89b 100644 --- a/mindspore/core/ops/fusion/avg_pool_fusion.cc +++ b/mindspore/core/ops/fusion/avg_pool_fusion.cc @@ -15,6 +15,9 @@ */ #include "ops/fusion/avg_pool_fusion.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -31,11 +34,11 @@ void AvgPoolFusion::Init(const std::vector &kernel_size, const std::vec this->set_activation_type(activation_type); } -void AvgPoolFusion::set_global(const bool global) { (void)AddAttr(kGlobal, MakeValue(global)); } +void AvgPoolFusion::set_global(const bool global) { (void)AddAttr(kGlobal, api::MakeValue(global)); } void AvgPoolFusion::set_activation_type(ActivationType activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } bool AvgPoolFusion::get_global() const { @@ -50,6 +53,7 @@ ActivationType AvgPoolFusion::get_activation_type() const { return ActivationType(GetValue(value_ptr)); } +MIND_API_BASE_IMPL(AvgPoolFusion, PrimitiveC, AvgPool); REGISTER_PRIMITIVE_C(kNameAvgPoolFusion, AvgPoolFusion); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/avg_pool_fusion.h b/mindspore/core/ops/fusion/avg_pool_fusion.h index 80122f454f..19dc965136 100644 --- a/mindspore/core/ops/fusion/avg_pool_fusion.h +++ b/mindspore/core/ops/fusion/avg_pool_fusion.h @@ -20,23 +20,18 @@ #include #include "ops/avg_pool.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAvgPoolFusion = "AvgPoolFusion"; /// \brief AvgPoolFusion defined AvgPool operator prototype of lite. -class MS_CORE_API AvgPoolFusion : public AvgPool { +class MIND_API AvgPoolFusion : public AvgPool { public: + MIND_API_BASE_MEMBER(AvgPoolFusion); /// \brief Constructor. AvgPoolFusion() : AvgPool(kNameAvgPoolFusion) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~AvgPoolFusion() = default; - - MS_DECLARE_PARENT(AvgPoolFusion, AvgPool); - /// \brief Method to init the op's attributes. /// /// \param[in] kernel_size Define the size of the kernel. @@ -75,8 +70,8 @@ class MS_CORE_API AvgPoolFusion : public AvgPool { ActivationType get_activation_type() const; }; -AbstractBasePtr AvgPoolFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AvgPoolFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.cc b/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.cc index 855fe3affd..be8dda0a74 100644 --- a/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.cc +++ b/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.cc @@ -18,9 +18,11 @@ #include "ops/fusion/conv2d_backprop_filter_fusion.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Conv2DBackpropFilterFusion, PrimitiveC, Conv2DBackpropFilter); void Conv2DBackpropFilterFusion::Init(const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &pad_list, const int64_t mode, const std::vector &stride, const std::vector &dilation, @@ -39,11 +41,11 @@ void Conv2DBackpropFilterFusion::Init(const int64_t out_channel, const std::vect void Conv2DBackpropFilterFusion::set_activation_type(const ActivationType activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } void Conv2DBackpropFilterFusion::set_in_channel(const int64_t in_channel) { - (void)this->AddAttr(kInChannel, MakeValue(in_channel)); + (void)this->AddAttr(kInChannel, api::MakeValue(in_channel)); } ActivationType Conv2DBackpropFilterFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.h b/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.h index ab6462b827..bf99bc12fb 100644 --- a/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.h +++ b/mindspore/core/ops/fusion/conv2d_backprop_filter_fusion.h @@ -20,25 +20,20 @@ #include #include "ops/grad/conv2d_backprop_filter.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameConv2DBackpropFilterFusion = "Conv2DBackpropFilterFusion"; /// \brief Conv2DBackpropFilterFusion defined Conv2DBackpropFilter operator prototype of lite. -class MS_CORE_API Conv2DBackpropFilterFusion : public Conv2DBackpropFilter { +class MIND_API Conv2DBackpropFilterFusion : public Conv2DBackpropFilter { public: + MIND_API_BASE_MEMBER(Conv2DBackpropFilterFusion); /// \brief Constructor. Conv2DBackpropFilterFusion() : Conv2DBackpropFilter(kNameConv2DBackpropFilterFusion) { InitIOName({"out_backprop", "input", "filter_sizes"}, {"output"}); } - /// \brief Destructor. - ~Conv2DBackpropFilterFusion() = default; - - MS_DECLARE_PARENT(Conv2DBackpropFilterFusion, Conv2DBackpropFilter); - /// \brief Method to init the op's attributes. /// /// \param[in] out_channel Define the number of output channel. diff --git a/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.cc b/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.cc index 32cea217f4..79c1af3295 100644 --- a/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.cc +++ b/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.cc @@ -19,9 +19,11 @@ #include "ops/fusion/conv2d_backprop_input_fusion.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Conv2DBackpropInputFusion, PrimitiveC, Conv2DBackpropInput); void Conv2DBackpropInputFusion::Init(int64_t in_channel, int64_t out_channel, const std::vector &kernel_size, int64_t mode, const PadMode &pad_mode, const std::vector &pad, const std::vector &stride, const std::vector &dilation, @@ -42,12 +44,12 @@ void Conv2DBackpropInputFusion::Init(int64_t in_channel, int64_t out_channel, co } void Conv2DBackpropInputFusion::set_in_channel(int64_t in_channel) { - (void)this->AddAttr(kInChannel, MakeValue(in_channel)); + (void)this->AddAttr(kInChannel, api::MakeValue(in_channel)); } void Conv2DBackpropInputFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } int64_t Conv2DBackpropInputFusion::get_in_channel() const { auto value_ptr = GetAttr(kInChannel); diff --git a/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.h b/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.h index 309fb16943..1eda4addaa 100644 --- a/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.h +++ b/mindspore/core/ops/fusion/conv2d_backprop_input_fusion.h @@ -18,23 +18,18 @@ #define MINDSPORE_CORE_OPS_CONV2D_BACKPROP_INPUT_FUSION_H_ #include #include "ops/grad/conv2d_backprop_input.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameConv2DBackpropInputFusion = "Conv2DBackpropInputFusion"; /// \brief Conv2DBackpropInputFusion defined Conv2DBackpropInput operator prototype of lite. -class MS_CORE_API Conv2DBackpropInputFusion : public Conv2DBackpropInput { +class MIND_API Conv2DBackpropInputFusion : public Conv2DBackpropInput { public: + MIND_API_BASE_MEMBER(Conv2DBackpropInputFusion); /// \brief Constructor. Conv2DBackpropInputFusion() : Conv2DBackpropInput(kNameConv2DBackpropInputFusion) {} - /// \brief Destructor. - ~Conv2DBackpropInputFusion() = default; - - MS_DECLARE_PARENT(Conv2DBackpropInputFusion, Conv2DBackpropInput); - /// \brief Method to init the op's attributes. /// /// \param[in] in_channel Define the number of input channel. diff --git a/mindspore/core/ops/fusion/conv2d_fusion.cc b/mindspore/core/ops/fusion/conv2d_fusion.cc index 9e04209df7..6be0e2d571 100644 --- a/mindspore/core/ops/fusion/conv2d_fusion.cc +++ b/mindspore/core/ops/fusion/conv2d_fusion.cc @@ -18,9 +18,11 @@ #include "ops/fusion/conv2d_fusion.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Conv2DFusion, PrimitiveC, Conv2D); void Conv2DFusion::Init(int64_t in_channel, int64_t out_channel, const std::vector &kernel_size, int64_t mode, const PadMode &pad_mode, const std::vector &pad, const std::vector &stride, const std::vector &dilation, int64_t group, const Format &format, @@ -38,13 +40,15 @@ void Conv2DFusion::Init(int64_t in_channel, int64_t out_channel, const std::vect this->set_pad_list(pad_list); this->set_activation_type(activation_type); } -void Conv2DFusion::set_in_channel(const int64_t in_channel) { (void)this->AddAttr(kInChannel, MakeValue(in_channel)); } +void Conv2DFusion::set_in_channel(const int64_t in_channel) { + (void)this->AddAttr(kInChannel, api::MakeValue(in_channel)); +} void Conv2DFusion::set_pad_list(const std::vector &pad_list) { - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } void Conv2DFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } int64_t Conv2DFusion::get_in_channel() const { auto value_ptr = GetAttr(kInChannel); diff --git a/mindspore/core/ops/fusion/conv2d_fusion.h b/mindspore/core/ops/fusion/conv2d_fusion.h index b8539507bb..9a4d87efff 100644 --- a/mindspore/core/ops/fusion/conv2d_fusion.h +++ b/mindspore/core/ops/fusion/conv2d_fusion.h @@ -19,23 +19,18 @@ #include #include "ops/conv2d.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameConv2DFusion = "Conv2DFusion"; /// \brief Conv2DFusion defined Conv2D operator prototype of lite. -class MS_CORE_API Conv2DFusion : public Conv2D { +class MIND_API Conv2DFusion : public Conv2D { public: + MIND_API_BASE_MEMBER(Conv2DFusion); /// \brief Constructor. Conv2DFusion() : Conv2D(kNameConv2DFusion) {} - /// \brief Destructor. - ~Conv2DFusion() = default; - - MS_DECLARE_PARENT(Conv2DFusion, Conv2D); - /// \brief Method to init the op's attributes. /// /// \param[in] in_channel Define the number of input channel. diff --git a/mindspore/core/ops/fusion/conv2d_transpose_fusion.cc b/mindspore/core/ops/fusion/conv2d_transpose_fusion.cc index f7ddcd8797..f708bd1cb2 100644 --- a/mindspore/core/ops/fusion/conv2d_transpose_fusion.cc +++ b/mindspore/core/ops/fusion/conv2d_transpose_fusion.cc @@ -16,9 +16,12 @@ #include "ops/fusion/conv2d_transpose_fusion.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Conv2dTransposeFusion, PrimitiveC, Conv2DTranspose); void Conv2dTransposeFusion::Init(int64_t in_channel, int64_t out_channel, const std::vector &kernel_size, int64_t mode, const PadMode &pad_mode, const std::vector &pad, const std::vector &stride, const std::vector &dilation, @@ -45,7 +48,7 @@ void Conv2dTransposeFusion::set_kernel_size(const std::vector &kernel_s for (int64_t item : kernel_size) { (void)CheckAndConvertUtils::CheckInteger(kKernelSize, item, kGreaterEqual, 1, name()); } - (void)AddAttr(kKernelSize, MakeValue(kernel_size)); + (void)AddAttr(kKernelSize, api::MakeValue(kernel_size)); } void Conv2dTransposeFusion::set_dilation(const std::vector &dilation) { @@ -54,7 +57,7 @@ void Conv2dTransposeFusion::set_dilation(const std::vector &dilation) { for (int64_t item : dilation) { (void)CheckAndConvertUtils::CheckInteger(kDilation, item, kGreaterEqual, 1, name()); } - (void)AddAttr(kDilation, MakeValue(dilation)); + (void)AddAttr(kDilation, api::MakeValue(dilation)); } void Conv2dTransposeFusion::set_output_paddings(const std::vector &output_paddings) { @@ -63,12 +66,12 @@ void Conv2dTransposeFusion::set_output_paddings(const std::vector &outp for (int64_t item : output_paddings) { (void)CheckAndConvertUtils::CheckInteger(kOutputPaddings, item, kGreaterEqual, 0, name()); } - (void)AddAttr(kOutputPaddings, MakeValue(output_paddings)); + (void)AddAttr(kOutputPaddings, api::MakeValue(output_paddings)); } void Conv2dTransposeFusion::set_activation_type(ActivationType activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } std::vector Conv2dTransposeFusion::get_output_paddings() const { diff --git a/mindspore/core/ops/fusion/conv2d_transpose_fusion.h b/mindspore/core/ops/fusion/conv2d_transpose_fusion.h index b6291aa579..dcaa37472b 100644 --- a/mindspore/core/ops/fusion/conv2d_transpose_fusion.h +++ b/mindspore/core/ops/fusion/conv2d_transpose_fusion.h @@ -19,25 +19,20 @@ #include #include "ops/conv2d_transpose.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameConv2dTransposeFusion = "Conv2dTransposeFusion"; /// \brief Conv2dTransposeFusion defined Conv2dTranspose operator prototype of lite. -class MS_CORE_API Conv2dTransposeFusion : public Conv2DTranspose { +class MIND_API Conv2dTransposeFusion : public Conv2DTranspose { public: + MIND_API_BASE_MEMBER(Conv2dTransposeFusion); /// \brief Constructor. Conv2dTransposeFusion() : Conv2DTranspose(kNameConv2dTransposeFusion) { InitIOName({"out_backprop", "filter", "input_sizes"}, {"output"}); } - /// \brief Destructor. - ~Conv2dTransposeFusion() = default; - - MS_DECLARE_PARENT(Conv2dTransposeFusion, Conv2DTranspose); - /// \brief Method to init the op's attributes. /// /// \param[in] in_channel Define the number of input channel. diff --git a/mindspore/core/ops/fusion/div_fusion.cc b/mindspore/core/ops/fusion/div_fusion.cc index 6b08f6152f..33330c3b0d 100644 --- a/mindspore/core/ops/fusion/div_fusion.cc +++ b/mindspore/core/ops/fusion/div_fusion.cc @@ -17,14 +17,16 @@ #include "ops/fusion/div_fusion.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(DivFusion, PrimitiveC, Div); void DivFusion::Init(const ActivationType &activation_type) { this->set_activation_type(activation_type); } void DivFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType DivFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/div_fusion.h b/mindspore/core/ops/fusion/div_fusion.h index 9aa5ae212e..401376958e 100644 --- a/mindspore/core/ops/fusion/div_fusion.h +++ b/mindspore/core/ops/fusion/div_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_DIV_FUSION_H_ #define MINDSPORE_CORE_OPS_DIV_FUSION_H_ #include "ops/div.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDivFusion = "DivFusion"; /// \brief DivFusion defined Div operator prototype of lite. -class MS_CORE_API DivFusion : public Div { +class MIND_API DivFusion : public Div { public: + MIND_API_BASE_MEMBER(DivFusion); /// \brief Constructor. DivFusion() : Div(kNameDivFusion) {} - /// \brief Destructor. - ~DivFusion() = default; - - MS_DECLARE_PARENT(DivFusion, Div); - /// \brief Method to init the op's attributes. /// /// \param[in] activation_type Define the activation type. diff --git a/mindspore/core/ops/fusion/embedding_lookup_fusion.cc b/mindspore/core/ops/fusion/embedding_lookup_fusion.cc index 4eee54ff45..e7c27d1603 100644 --- a/mindspore/core/ops/fusion/embedding_lookup_fusion.cc +++ b/mindspore/core/ops/fusion/embedding_lookup_fusion.cc @@ -17,10 +17,14 @@ #include "ops/fusion/embedding_lookup_fusion.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void EmbeddingLookupFusion::set_max_norm(const float max_norm) { (void)this->AddAttr(kMaxNorm, MakeValue(max_norm)); } +MIND_API_BASE_IMPL(EmbeddingLookupFusion, PrimitiveC, BaseOperator); +void EmbeddingLookupFusion::set_max_norm(const float max_norm) { + (void)this->AddAttr(kMaxNorm, api::MakeValue(max_norm)); +} float EmbeddingLookupFusion::get_max_norm() const { auto value_ptr = GetAttr(kMaxNorm); MS_EXCEPTION_IF_NULL(value_ptr); diff --git a/mindspore/core/ops/fusion/embedding_lookup_fusion.h b/mindspore/core/ops/fusion/embedding_lookup_fusion.h index 1288e41743..6199b8b66f 100644 --- a/mindspore/core/ops/fusion/embedding_lookup_fusion.h +++ b/mindspore/core/ops/fusion/embedding_lookup_fusion.h @@ -20,26 +20,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameEmbeddingLookupFusion = "EmbeddingLookupFusion"; /// \brief EmbeddingLookupFusion defined EmbeddingLookup operator prototype of lite. -class MS_CORE_API EmbeddingLookupFusion : public PrimitiveC { +class MIND_API EmbeddingLookupFusion : public BaseOperator { public: + MIND_API_BASE_MEMBER(EmbeddingLookupFusion); /// \brief Constructor. - EmbeddingLookupFusion() : PrimitiveC(kNameEmbeddingLookupFusion) { + EmbeddingLookupFusion() : BaseOperator(kNameEmbeddingLookupFusion) { InitIOName({"params", "indices", "offset"}, {"output"}); } - /// \brief Destructor. - ~EmbeddingLookupFusion() = default; - - MS_DECLARE_PARENT(EmbeddingLookupFusion, PrimitiveC); - /// \brief Method to init the op's attributes. /// /// \param[in] max_norm Define the max l2-norm value of each embedding. Each embedding will be clip if l2-norm is diff --git a/mindspore/core/ops/fusion/exp_fusion.cc b/mindspore/core/ops/fusion/exp_fusion.cc index 9549ed7154..da3e9c1216 100644 --- a/mindspore/core/ops/fusion/exp_fusion.cc +++ b/mindspore/core/ops/fusion/exp_fusion.cc @@ -20,20 +20,22 @@ #include #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ExpFusion, PrimitiveC, Exp); void ExpFusion::Init(const float base, const float scale, const float shift) { this->set_base(base); this->set_scale(scale); this->set_shift(shift); } -void ExpFusion::set_base(const float base) { (void)this->AddAttr(kBase, MakeValue(base)); } +void ExpFusion::set_base(const float base) { (void)this->AddAttr(kBase, api::MakeValue(base)); } -void ExpFusion::set_scale(const float scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } +void ExpFusion::set_scale(const float scale) { (void)this->AddAttr(kScale, api::MakeValue(scale)); } -void ExpFusion::set_shift(const float shift) { (void)this->AddAttr(kShift, MakeValue(shift)); } +void ExpFusion::set_shift(const float shift) { (void)this->AddAttr(kShift, api::MakeValue(shift)); } float ExpFusion::get_base() const { auto value_ptr = GetAttr(kBase); diff --git a/mindspore/core/ops/fusion/exp_fusion.h b/mindspore/core/ops/fusion/exp_fusion.h index 2230f775f8..8ccbf8bb73 100644 --- a/mindspore/core/ops/fusion/exp_fusion.h +++ b/mindspore/core/ops/fusion/exp_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_EXP_FUSION_H_ #define MINDSPORE_CORE_OPS_EXP_FUSION_H_ #include "ops/exp.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameExpFusion = "ExpFusion"; /// \brief ExpFusion defined Exp operator prototype of lite. -class MS_CORE_API ExpFusion : public Exp { +class MIND_API ExpFusion : public Exp { public: + MIND_API_BASE_MEMBER(ExpFusion); /// \brief Constructor. ExpFusion() : Exp(kNameExpFusion) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~ExpFusion() = default; - - MS_DECLARE_PARENT(ExpFusion, Exp); - /// \brief Method to init the op's attributes. /// /// \param[in] base Define a base number. If base is -1, it represents e. In addition to this, base must be larger diff --git a/mindspore/core/ops/fusion/full_connection.cc b/mindspore/core/ops/fusion/full_connection.cc index 61a000985b..a2be4f5615 100644 --- a/mindspore/core/ops/fusion/full_connection.cc +++ b/mindspore/core/ops/fusion/full_connection.cc @@ -17,10 +17,13 @@ #include "ops/fusion/full_connection.h" #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void FullConnection::set_has_bias(const bool has_bias) { (void)this->AddAttr(kHasBias, MakeValue(has_bias)); } +MIND_API_BASE_IMPL(FullConnection, PrimitiveC, BaseOperator); +void FullConnection::set_has_bias(const bool has_bias) { (void)this->AddAttr(kHasBias, api::MakeValue(has_bias)); } bool FullConnection::get_has_bias() const { auto value_ptr = GetAttr(kHasBias); @@ -28,14 +31,14 @@ bool FullConnection::get_has_bias() const { return GetValue(value_ptr); } -void FullConnection::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void FullConnection::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t FullConnection::get_axis() const { auto value_ptr = GetAttr(kAxis); MS_EXCEPTION_IF_NULL(value_ptr); return GetValue(value_ptr); } -void FullConnection::set_use_axis(const bool use_axis) { (void)this->AddAttr(kUseAxis, MakeValue(use_axis)); } +void FullConnection::set_use_axis(const bool use_axis) { (void)this->AddAttr(kUseAxis, api::MakeValue(use_axis)); } bool FullConnection::get_use_axis() const { auto value_ptr = GetAttr(kUseAxis); MS_EXCEPTION_IF_NULL(value_ptr); @@ -44,7 +47,7 @@ bool FullConnection::get_use_axis() const { void FullConnection::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType FullConnection::get_activation_type() const { auto value_ptr = GetAttr(kActivationType); diff --git a/mindspore/core/ops/fusion/full_connection.h b/mindspore/core/ops/fusion/full_connection.h index e050a6c4e8..408f55340b 100644 --- a/mindspore/core/ops/fusion/full_connection.h +++ b/mindspore/core/ops/fusion/full_connection.h @@ -20,23 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFullConnection = "FullConnection"; /// \brief FullConnection defined FullConnection operator prototype of lite. -class MS_CORE_API FullConnection : public PrimitiveC { +class MIND_API FullConnection : public BaseOperator { public: + MIND_API_BASE_MEMBER(FullConnection); /// \brief Constructor. - FullConnection() : PrimitiveC(kNameFullConnection) { InitIOName({"x1", "x2", "b"}, {"output"}); } - - /// \brief Destructor. - ~FullConnection() = default; - - MS_DECLARE_PARENT(FullConnection, PrimitiveC); + FullConnection() : BaseOperator(kNameFullConnection) { InitIOName({"x1", "x2", "b"}, {"output"}); } /// \brief Method to init the op's attributes. /// @@ -86,8 +81,8 @@ class MS_CORE_API FullConnection : public PrimitiveC { /// \return activation type. ActivationType get_activation_type() const; }; -AbstractBasePtr FullConnectionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FullConnectionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/l2_normalize_fusion.cc b/mindspore/core/ops/fusion/l2_normalize_fusion.cc index 9ec3c0fc16..933bf275d7 100644 --- a/mindspore/core/ops/fusion/l2_normalize_fusion.cc +++ b/mindspore/core/ops/fusion/l2_normalize_fusion.cc @@ -17,9 +17,11 @@ #include "ops/fusion/l2_normalize_fusion.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(L2NormalizeFusion, PrimitiveC, L2Normalize); void L2NormalizeFusion::Init(const std::vector &axis, const float epsilon, const ActivationType &activation_type) { this->set_axis(axis); @@ -29,7 +31,7 @@ void L2NormalizeFusion::Init(const std::vector &axis, const float epsil void L2NormalizeFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType L2NormalizeFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/l2_normalize_fusion.h b/mindspore/core/ops/fusion/l2_normalize_fusion.h index 877045c735..90f8fc0c54 100644 --- a/mindspore/core/ops/fusion/l2_normalize_fusion.h +++ b/mindspore/core/ops/fusion/l2_normalize_fusion.h @@ -19,23 +19,18 @@ #include #include "ops/l2_normalize.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameL2NormalizeFusion = "L2NormalizeFusion"; /// \brief L2NormalizeFusion defined L2Normalize operator prototype of lite. -class MS_CORE_API L2NormalizeFusion : public L2Normalize { +class MIND_API L2NormalizeFusion : public L2Normalize { public: + MIND_API_BASE_MEMBER(L2NormalizeFusion); /// \brief Constructor. L2NormalizeFusion() : L2Normalize(kNameL2NormalizeFusion) {} - /// \brief Destructor. - ~L2NormalizeFusion() = default; - - MS_DECLARE_PARENT(L2NormalizeFusion, L2Normalize); - /// \brief Method to init the op's attributes. /// /// \param[in] axis Define a axis that the normalization is done along with. diff --git a/mindspore/core/ops/fusion/layer_norm_fusion.cc b/mindspore/core/ops/fusion/layer_norm_fusion.cc index 8ff0e9abf4..1551ee701a 100644 --- a/mindspore/core/ops/fusion/layer_norm_fusion.cc +++ b/mindspore/core/ops/fusion/layer_norm_fusion.cc @@ -15,9 +15,12 @@ */ #include "ops/fusion/layer_norm_fusion.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(LayerNormFusion, PrimitiveC, LayerNorm); void LayerNormFusion::Init(const int64_t begin_norm_axis, const int64_t begin_params_axis, const float epsilon, const bool elementwise_affine) { this->set_begin_norm_axis(begin_norm_axis); @@ -27,7 +30,7 @@ void LayerNormFusion::Init(const int64_t begin_norm_axis, const int64_t begin_pa } void LayerNormFusion::set_elementwise_affine(const bool elementwise_affine) { - (void)AddAttr(kElementwiseAffine, MakeValue(elementwise_affine)); + (void)AddAttr(kElementwiseAffine, api::MakeValue(elementwise_affine)); } bool LayerNormFusion::get_elementwise_affine() const { diff --git a/mindspore/core/ops/fusion/layer_norm_fusion.h b/mindspore/core/ops/fusion/layer_norm_fusion.h index 498113ac4e..07a8b25bbb 100644 --- a/mindspore/core/ops/fusion/layer_norm_fusion.h +++ b/mindspore/core/ops/fusion/layer_norm_fusion.h @@ -20,23 +20,18 @@ #include #include "ops/layer_norm.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLayerNormFusion = "LayerNormFusion"; /// \brief LayerNormFusion defined LayerNorm operator prototype of lite. -class MS_CORE_API LayerNormFusion : public LayerNorm { +class MIND_API LayerNormFusion : public LayerNorm { public: + MIND_API_BASE_MEMBER(LayerNormFusion); /// \brief Constructor. LayerNormFusion() : LayerNorm(kNameLayerNormFusion) {} - /// \brief Destructor. - ~LayerNormFusion() = default; - - MS_DECLARE_PARENT(LayerNormFusion, LayerNorm); - /// \brief Method to init the op's attributes. /// /// \param[in] begin_norm_axis Define the first normalization dimension of input. @@ -57,8 +52,8 @@ class MS_CORE_API LayerNormFusion : public LayerNorm { bool get_elementwise_affine() const; }; -AbstractBasePtr LayerNormFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/mat_mul_fusion.cc b/mindspore/core/ops/fusion/mat_mul_fusion.cc index 5c36eed7de..8f5660ffcd 100644 --- a/mindspore/core/ops/fusion/mat_mul_fusion.cc +++ b/mindspore/core/ops/fusion/mat_mul_fusion.cc @@ -18,9 +18,13 @@ #include #include #include "ops/fusion/mat_mul_fusion.h" +#include "ops/op_utils.h" +#include "abstract/dshape.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(MatMulFusion, PrimitiveC, MatMul); void MatMulFusion::Init(bool transpose_a, bool transpose_b, const ActivationType &activation_type) { set_transpose_a(transpose_a); set_transpose_b(transpose_b); @@ -29,7 +33,7 @@ void MatMulFusion::Init(bool transpose_a, bool transpose_b, const ActivationType void MatMulFusion::set_activation_type(const ActivationType activation_type) { int64_t act = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(act)); + (void)this->AddAttr(kActivationType, api::MakeValue(act)); } ActivationType MatMulFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/mat_mul_fusion.h b/mindspore/core/ops/fusion/mat_mul_fusion.h index f7eb06a209..81fbd92069 100644 --- a/mindspore/core/ops/fusion/mat_mul_fusion.h +++ b/mindspore/core/ops/fusion/mat_mul_fusion.h @@ -20,23 +20,18 @@ #include #include "ops/mat_mul.h" -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "abstract/dshape.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMatMulFusion = "MatMulFusion"; /// \brief Multiplies matrix a and matrix b. Refer to Python API @ref mindspore.ops.MatMul for more details. -class MS_CORE_API MatMulFusion : public MatMul { +class MIND_API MatMulFusion : public MatMul { public: + MIND_API_BASE_MEMBER(MatMulFusion); /// \brief Constructor. MatMulFusion() : MatMul(kNameMatMulFusion) {} - /// \brief Destructor. - ~MatMulFusion() = default; - MS_DECLARE_PARENT(MatMulFusion, MatMul); /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.MatMulFusion for the inputs. void Init(bool transpose_a = false, bool transpose_b = false, const ActivationType &activation_type = NO_ACTIVATION); /// \brief Method to set activation type. diff --git a/mindspore/core/ops/fusion/max_pool_fusion.cc b/mindspore/core/ops/fusion/max_pool_fusion.cc index 2531661105..7db68a2efa 100644 --- a/mindspore/core/ops/fusion/max_pool_fusion.cc +++ b/mindspore/core/ops/fusion/max_pool_fusion.cc @@ -15,6 +15,9 @@ */ #include "ops/fusion/max_pool_fusion.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -31,11 +34,11 @@ void MaxPoolFusion::Init(const std::vector &kernel_size, const std::vec this->set_activation_type(activation_type); } -void MaxPoolFusion::set_global(const bool global) { (void)AddAttr(kGlobal, MakeValue(global)); } +void MaxPoolFusion::set_global(const bool global) { (void)AddAttr(kGlobal, api::MakeValue(global)); } void MaxPoolFusion::set_activation_type(ActivationType activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } bool MaxPoolFusion::get_global() const { @@ -50,6 +53,7 @@ ActivationType MaxPoolFusion::get_activation_type() const { return ActivationType(GetValue(value_ptr)); } +MIND_API_BASE_IMPL(MaxPoolFusion, PrimitiveC, MaxPool); REGISTER_PRIMITIVE_C(kNameMaxPoolFusion, MaxPoolFusion); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/max_pool_fusion.h b/mindspore/core/ops/fusion/max_pool_fusion.h index b2ef956812..59b411b5c7 100644 --- a/mindspore/core/ops/fusion/max_pool_fusion.h +++ b/mindspore/core/ops/fusion/max_pool_fusion.h @@ -20,23 +20,18 @@ #include #include "ops/max_pool.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMaxPoolFusion = "MaxPoolFusion"; /// \brief MaxPoolFusion defined MaxPool operator prototype of lite. -class MS_CORE_API MaxPoolFusion : public MaxPool { +class MIND_API MaxPoolFusion : public MaxPool { public: + MIND_API_BASE_MEMBER(MaxPoolFusion); /// \brief Constructor. MaxPoolFusion() : MaxPool(kNameMaxPoolFusion) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~MaxPoolFusion() = default; - - MS_DECLARE_PARENT(MaxPoolFusion, MaxPool); - /// \brief Method to init the op's attributes. /// /// \param[in] kernel_size Define the size of the kernel. @@ -75,8 +70,8 @@ class MS_CORE_API MaxPoolFusion : public MaxPool { ActivationType get_activation_type() const; }; -AbstractBasePtr MaxPoolFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MaxPoolFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/mul_fusion.cc b/mindspore/core/ops/fusion/mul_fusion.cc index 68028b53ae..2a871b847f 100644 --- a/mindspore/core/ops/fusion/mul_fusion.cc +++ b/mindspore/core/ops/fusion/mul_fusion.cc @@ -20,12 +20,14 @@ #include #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(MulFusion, PrimitiveC, Mul); void MulFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType MulFusion::get_activation_type() const { auto value_ptr = GetAttr(kActivationType); diff --git a/mindspore/core/ops/fusion/mul_fusion.h b/mindspore/core/ops/fusion/mul_fusion.h index b771b497a0..4b459bde46 100644 --- a/mindspore/core/ops/fusion/mul_fusion.h +++ b/mindspore/core/ops/fusion/mul_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_MUL_FUSION_H_ #define MINDSPORE_CORE_OPS_MUL_FUSION_H_ #include "ops/mul.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMulFusion = "MulFusion"; /// \brief MulFusion defined Mul operator prototype of lite. -class MS_CORE_API MulFusion : public Mul { +class MIND_API MulFusion : public Mul { public: + MIND_API_BASE_MEMBER(MulFusion); /// \brief Constructor. MulFusion() : Mul(kNameMulFusion) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~MulFusion() = default; - - MS_DECLARE_PARENT(MulFusion, Mul); - /// \brief Method to init the op's attributes. /// /// \param[in] activation_type Define the activation type. diff --git a/mindspore/core/ops/fusion/pad_fusion.cc b/mindspore/core/ops/fusion/pad_fusion.cc index fc556343d1..a766462cc1 100644 --- a/mindspore/core/ops/fusion/pad_fusion.cc +++ b/mindspore/core/ops/fusion/pad_fusion.cc @@ -20,9 +20,11 @@ #include #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PadFusion, PrimitiveC, Pad); void PadFusion::Init(const PaddingMode &padding_mode, const float constant_value) { this->set_padding_mode(padding_mode); this->set_constant_value(constant_value); @@ -30,11 +32,11 @@ void PadFusion::Init(const PaddingMode &padding_mode, const float constant_value void PadFusion::set_padding_mode(const PaddingMode &padding_mode) { int64_t swi = padding_mode; - (void)this->AddAttr(kPaddingMode, MakeValue(swi)); + (void)this->AddAttr(kPaddingMode, api::MakeValue(swi)); } void PadFusion::set_constant_value(const float constant_value) { - (void)this->AddAttr(kConstantValue, MakeValue(constant_value)); + (void)this->AddAttr(kConstantValue, api::MakeValue(constant_value)); } PaddingMode PadFusion::get_padding_mode() const { diff --git a/mindspore/core/ops/fusion/pad_fusion.h b/mindspore/core/ops/fusion/pad_fusion.h index 216dd9c7da..eddf54f3d5 100644 --- a/mindspore/core/ops/fusion/pad_fusion.h +++ b/mindspore/core/ops/fusion/pad_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_PAD_FUSION_H_ #define MINDSPORE_CORE_OPS_PAD_FUSION_H_ #include "ops/pad.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePadFusion = "PadFusion"; /// \brief PadFusion defined Pad operator prototype of lite. -class MS_CORE_API PadFusion : public Pad { +class MIND_API PadFusion : public Pad { public: + MIND_API_BASE_MEMBER(PadFusion); /// \brief Constructor. PadFusion() : Pad(kNamePadFusion) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~PadFusion() = default; - - MS_DECLARE_PARENT(PadFusion, Pad); - /// \brief Method to init the op's attributes. /// /// \param[in] padding_mode Define the padding mode. diff --git a/mindspore/core/ops/fusion/partial_fusion.cc b/mindspore/core/ops/fusion/partial_fusion.cc index 14108dcde0..90810dea57 100644 --- a/mindspore/core/ops/fusion/partial_fusion.cc +++ b/mindspore/core/ops/fusion/partial_fusion.cc @@ -16,12 +16,14 @@ #include "ops/fusion/partial_fusion.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PartialFusion, PrimitiveC, BaseOperator); void PartialFusion::Init(const int64_t sub_graph_index) { this->set_sub_graph_index(sub_graph_index); } void PartialFusion::set_sub_graph_index(const int64_t sub_graph_index) { - (void)this->AddAttr(kSubGraphIndex, MakeValue(sub_graph_index)); + (void)this->AddAttr(kSubGraphIndex, api::MakeValue(sub_graph_index)); } int64_t PartialFusion::get_sub_graph_index() const { auto value_ptr = GetAttr(kSubGraphIndex); diff --git a/mindspore/core/ops/fusion/partial_fusion.h b/mindspore/core/ops/fusion/partial_fusion.h index 0e53839dd9..7a60f9ded7 100644 --- a/mindspore/core/ops/fusion/partial_fusion.h +++ b/mindspore/core/ops/fusion/partial_fusion.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_PARTIAL_FUSION_H_ #define MINDSPORE_CORE_OPS_PARTIAL_FUSION_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePartialFusion = "PartialFusion"; /// \brief PartialFusion defined Partial operator prototype of lite. -class MS_CORE_API PartialFusion : public PrimitiveC { +class MIND_API PartialFusion : public BaseOperator { public: + MIND_API_BASE_MEMBER(PartialFusion); /// \brief Constructor. - PartialFusion() : PrimitiveC(kNamePartialFusion) {} - - /// \brief Destructor. - ~PartialFusion() = default; - - MS_DECLARE_PARENT(PartialFusion, PrimitiveC); + PartialFusion() : BaseOperator(kNamePartialFusion) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/fusion/pow_fusion.cc b/mindspore/core/ops/fusion/pow_fusion.cc index 03a39c8d48..015ef488dd 100644 --- a/mindspore/core/ops/fusion/pow_fusion.cc +++ b/mindspore/core/ops/fusion/pow_fusion.cc @@ -20,6 +20,8 @@ #include #include "ops/fusion/pow_fusion.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -28,12 +30,13 @@ void PowFusion::Init(const float &scale, const float &shift) { this->set_shift(shift); } -void PowFusion::set_scale(const float &scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } -void PowFusion::set_shift(const float &shift) { (void)this->AddAttr(kShift, MakeValue(shift)); } +void PowFusion::set_scale(const float &scale) { (void)this->AddAttr(kScale, api::MakeValue(scale)); } +void PowFusion::set_shift(const float &shift) { (void)this->AddAttr(kShift, api::MakeValue(shift)); } float PowFusion::get_scale() const { return GetValue(GetAttr(kScale)); } float PowFusion::get_shift() const { return GetValue(GetAttr(kShift)); } +MIND_API_BASE_IMPL(PowFusion, PrimitiveC, Pow); REGISTER_PRIMITIVE_C(kNamePowFusion, PowFusion); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/pow_fusion.h b/mindspore/core/ops/fusion/pow_fusion.h index e8591334fe..01d58ff2b3 100644 --- a/mindspore/core/ops/fusion/pow_fusion.h +++ b/mindspore/core/ops/fusion/pow_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_POW_FUSION_H_ #define MINDSPORE_CORE_OPS_POW_FUSION_H_ #include "ops/pow.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePowFusion = "PowFusion"; /// \brief PowFusion defined Pow operator prototype of lite. -class MS_CORE_API PowFusion : public Pow { +class MIND_API PowFusion : public Pow { public: + MIND_API_BASE_MEMBER(PowFusion); /// \brief Constructor. PowFusion() : Pow(kNamePowFusion) {} - /// \brief Destructor. - ~PowFusion() = default; - - MS_DECLARE_PARENT(PowFusion, Pow); - /// \brief Method to init the op's attributes. /// /// \param[in] scale Define a size factor applied to input. diff --git a/mindspore/core/ops/fusion/prelu_fusion.cc b/mindspore/core/ops/fusion/prelu_fusion.cc index 10ea5b5e66..145eeec2ac 100644 --- a/mindspore/core/ops/fusion/prelu_fusion.cc +++ b/mindspore/core/ops/fusion/prelu_fusion.cc @@ -20,19 +20,21 @@ #include #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PReLUFusion, PrimitiveC, PReLU); void PReLUFusion::Init(const bool channel_shared, const std::vector &slope) { this->set_channel_shared(channel_shared); this->set_slope(slope); } void PReLUFusion::set_channel_shared(const bool channel_shared) { - (void)this->AddAttr(kChannelShared, MakeValue(channel_shared)); + (void)this->AddAttr(kChannelShared, api::MakeValue(channel_shared)); } -void PReLUFusion::set_slope(const std::vector &slope) { (void)this->AddAttr(kSlope, MakeValue(slope)); } +void PReLUFusion::set_slope(const std::vector &slope) { (void)this->AddAttr(kSlope, api::MakeValue(slope)); } bool PReLUFusion::get_channel_shared() const { auto value_ptr = GetAttr(kChannelShared); diff --git a/mindspore/core/ops/fusion/prelu_fusion.h b/mindspore/core/ops/fusion/prelu_fusion.h index 20f3dd3666..cb8fc78b1b 100644 --- a/mindspore/core/ops/fusion/prelu_fusion.h +++ b/mindspore/core/ops/fusion/prelu_fusion.h @@ -19,23 +19,18 @@ #include #include "ops/prelu.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePReLUFusion = "PReLUFusion"; /// \brief PReLUFusion defined PReLU operator prototype of lite. -class MS_CORE_API PReLUFusion : public PReLU { +class MIND_API PReLUFusion : public PReLU { public: + MIND_API_BASE_MEMBER(PReLUFusion); /// \brief Constructor. PReLUFusion() : PReLU(kNamePReLUFusion) {} - /// \brief Destructor. - ~PReLUFusion() = default; - - MS_DECLARE_PARENT(PReLUFusion, PReLU); - /// \brief Method to init the op's attributes. /// /// \param[in] channel_shared Define a boolean value to indicate whether channel is shared or not. diff --git a/mindspore/core/ops/fusion/reduce_fusion.cc b/mindspore/core/ops/fusion/reduce_fusion.cc index 07df2dacbb..ff4ac82f5e 100644 --- a/mindspore/core/ops/fusion/reduce_fusion.cc +++ b/mindspore/core/ops/fusion/reduce_fusion.cc @@ -23,21 +23,23 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void ReduceFusion::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, MakeValue(keep_dims)); } +MIND_API_BASE_IMPL(ReduceFusion, PrimitiveC, Reduce); +void ReduceFusion::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, api::MakeValue(keep_dims)); } void ReduceFusion::set_mode(const ReduceMode mode) { int64_t swi = mode; - (void)this->AddAttr(kMode, MakeValue(swi)); + (void)this->AddAttr(kMode, api::MakeValue(swi)); } void ReduceFusion::set_reduce_to_end(const bool reduce_to_end) { - (void)this->AddAttr(kReduceToEnd, MakeValue(reduce_to_end)); + (void)this->AddAttr(kReduceToEnd, api::MakeValue(reduce_to_end)); } -void ReduceFusion::set_coeff(const float coeff) { (void)this->AddAttr(kCoeff, MakeValue(coeff)); } +void ReduceFusion::set_coeff(const float coeff) { (void)this->AddAttr(kCoeff, api::MakeValue(coeff)); } bool ReduceFusion::get_keep_dims() const { auto value_ptr = GetAttr(kKeepDims); diff --git a/mindspore/core/ops/fusion/reduce_fusion.h b/mindspore/core/ops/fusion/reduce_fusion.h index 9ce57c9d8d..9b9b8360f1 100644 --- a/mindspore/core/ops/fusion/reduce_fusion.h +++ b/mindspore/core/ops/fusion/reduce_fusion.h @@ -21,23 +21,18 @@ #include #include #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceFusion = "ReduceFusion"; /// \brief ReduceFusion defined Reduce operator prototype of lite. -class MS_CORE_API ReduceFusion : public Reduce { +class MIND_API ReduceFusion : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceFusion); /// \brief Constructor. ReduceFusion() : Reduce(kNameReduceFusion) {} - /// \brief Destructor. - ~ReduceFusion() = default; - - MS_DECLARE_PARENT(ReduceFusion, PrimitiveC); - /// \brief Method to init the op's attributes. /// /// \param[in] keep_dims Define a boolean value to indicate whether output dimension is kept or not. @@ -89,8 +84,8 @@ class MS_CORE_API ReduceFusion : public Reduce { /// \return a size factor applied to output. float get_coeff() const; }; -AbstractBasePtr ReduceFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReduceFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/scale_fusion.cc b/mindspore/core/ops/fusion/scale_fusion.cc index 6c572e8bf7..9fd8bb3170 100644 --- a/mindspore/core/ops/fusion/scale_fusion.cc +++ b/mindspore/core/ops/fusion/scale_fusion.cc @@ -18,9 +18,11 @@ #include #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ScaleFusion, PrimitiveC, Scale); void ScaleFusion::Init(const int64_t axis, const ActivationType &activation_type) { this->set_axis(axis); this->set_activation_type(activation_type); @@ -28,7 +30,7 @@ void ScaleFusion::Init(const int64_t axis, const ActivationType &activation_type void ScaleFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType ScaleFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/scale_fusion.h b/mindspore/core/ops/fusion/scale_fusion.h index d10653b0e2..0886873ece 100644 --- a/mindspore/core/ops/fusion/scale_fusion.h +++ b/mindspore/core/ops/fusion/scale_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_SCALE_FUSION_H_ #define MINDSPORE_CORE_OPS_SCALE_FUSION_H_ #include "ops/scale.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameScaleFusion = "ScaleFusion"; /// \brief ScaleFusion defined Scale operator prototype of lite. -class MS_CORE_API ScaleFusion : public Scale { +class MIND_API ScaleFusion : public Scale { public: + MIND_API_BASE_MEMBER(ScaleFusion); /// \brief Constructor. ScaleFusion() : Scale(kNameScaleFusion) {} - /// \brief Destructor. - ~ScaleFusion() = default; - - MS_DECLARE_PARENT(ScaleFusion, Scale); - /// \brief Method to init the op's attributes. /// /// \param[in] axis Define the first axis to do this operation. diff --git a/mindspore/core/ops/fusion/slice_fusion.cc b/mindspore/core/ops/fusion/slice_fusion.cc index dfb5be4bbb..9a8dd79a69 100644 --- a/mindspore/core/ops/fusion/slice_fusion.cc +++ b/mindspore/core/ops/fusion/slice_fusion.cc @@ -17,12 +17,15 @@ #include "ops/fusion/slice_fusion.h" #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SliceFusion, PrimitiveC, BaseOperator); void SliceFusion::Init(const std::vector &axes) { this->set_axes(axes); } -void SliceFusion::set_axes(const std::vector &axes) { (void)this->AddAttr(kAxes, MakeValue(axes)); } +void SliceFusion::set_axes(const std::vector &axes) { (void)this->AddAttr(kAxes, api::MakeValue(axes)); } std::vector SliceFusion::get_axes() const { auto value_ptr = GetAttr(kAxes); diff --git a/mindspore/core/ops/fusion/slice_fusion.h b/mindspore/core/ops/fusion/slice_fusion.h index 4fc471c688..584411f678 100644 --- a/mindspore/core/ops/fusion/slice_fusion.h +++ b/mindspore/core/ops/fusion/slice_fusion.h @@ -20,23 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSliceFusion = "SliceFusion"; /// \brief SliceFusion defined Slice operator prototype of lite. -class MS_CORE_API SliceFusion : public PrimitiveC { +class MIND_API SliceFusion : public BaseOperator { public: + MIND_API_BASE_MEMBER(SliceFusion); /// \brief Constructor. - SliceFusion() : PrimitiveC(kNameSliceFusion) { InitIOName({"x", "begin", "size"}, {"output"}); } - - /// \brief Destructor. - ~SliceFusion() = default; - - MS_DECLARE_PARENT(SliceFusion, PrimitiveC); + SliceFusion() : BaseOperator(kNameSliceFusion) { InitIOName({"x", "begin", "size"}, {"output"}); } /// \brief Method to init the op's attributes. /// @@ -53,8 +48,8 @@ class MS_CORE_API SliceFusion : public PrimitiveC { /// \return axes. std::vector get_axes() const; }; -AbstractBasePtr SliceFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SliceFusionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/fusion/sub_fusion.cc b/mindspore/core/ops/fusion/sub_fusion.cc index cb4ebba2f9..d82af96af1 100644 --- a/mindspore/core/ops/fusion/sub_fusion.cc +++ b/mindspore/core/ops/fusion/sub_fusion.cc @@ -17,14 +17,16 @@ #include "ops/fusion/sub_fusion.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SubFusion, PrimitiveC, Sub); void SubFusion::Init(const ActivationType &activation_type) { this->set_activation_type(activation_type); } void SubFusion::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType SubFusion::get_activation_type() const { diff --git a/mindspore/core/ops/fusion/sub_fusion.h b/mindspore/core/ops/fusion/sub_fusion.h index 1d8dc52394..1dd9767644 100644 --- a/mindspore/core/ops/fusion/sub_fusion.h +++ b/mindspore/core/ops/fusion/sub_fusion.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_SUB_FUSION_H_ #define MINDSPORE_CORE_OPS_SUB_FUSION_H_ #include "ops/sub.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSubFusion = "SubFusion"; /// \brief SubFusion defined Sub operator prototype of lite. -class MS_CORE_API SubFusion : public Sub { +class MIND_API SubFusion : public Sub { public: + MIND_API_BASE_MEMBER(SubFusion); /// \brief Constructor. SubFusion() : Sub(kNameSubFusion) {} - /// \brief Destructor. - ~SubFusion() = default; - - MS_DECLARE_PARENT(SubFusion, Sub); - /// \brief Method to init the op's attributes. /// /// \param[in] activation_type Define the activation type. diff --git a/mindspore/core/ops/fusion/tile_fusion.cc b/mindspore/core/ops/fusion/tile_fusion.cc index 69d8181c9a..ad24f44ad9 100644 --- a/mindspore/core/ops/fusion/tile_fusion.cc +++ b/mindspore/core/ops/fusion/tile_fusion.cc @@ -17,12 +17,14 @@ #include "ops/fusion/tile_fusion.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TileFusion, PrimitiveC, Tile); void TileFusion::Init(const std::vector &dims) { this->set_dims(dims); } -void TileFusion::set_dims(const std::vector &dims) { (void)this->AddAttr(kDims, MakeValue(dims)); } +void TileFusion::set_dims(const std::vector &dims) { (void)this->AddAttr(kDims, api::MakeValue(dims)); } std::vector TileFusion::get_dims() const { auto value_ptr = GetAttr(kDims); diff --git a/mindspore/core/ops/fusion/tile_fusion.h b/mindspore/core/ops/fusion/tile_fusion.h index 5e9dce5a25..f0ecefd108 100644 --- a/mindspore/core/ops/fusion/tile_fusion.h +++ b/mindspore/core/ops/fusion/tile_fusion.h @@ -19,23 +19,18 @@ #include #include "ops/tile.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTileFusion = "TileFusion"; /// \brief TileFusion defined Tile operator prototype of lite. -class MS_CORE_API TileFusion : public Tile { +class MIND_API TileFusion : public Tile { public: + MIND_API_BASE_MEMBER(TileFusion); /// \brief Constructor. TileFusion() : Tile(kNameTileFusion) {} - /// \brief Destructor. - ~TileFusion() = default; - - MS_DECLARE_PARENT(TileFusion, Tile); - /// \brief Method to init the op's attributes. /// /// \param[in] dims Define this operation will be performed on which axes. diff --git a/mindspore/core/ops/fusion/topk_fusion.cc b/mindspore/core/ops/fusion/topk_fusion.cc index 52d15362d5..492bf75ace 100644 --- a/mindspore/core/ops/fusion/topk_fusion.cc +++ b/mindspore/core/ops/fusion/topk_fusion.cc @@ -17,18 +17,20 @@ #include "ops/fusion/topk_fusion.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TopKFusion, PrimitiveC, TopK); void TopKFusion::Init(const bool sorted, const int64_t axis, const int64_t largest) { this->set_axis(axis); this->set_largest(largest); this->set_sorted(sorted); } -void TopKFusion::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void TopKFusion::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } -void TopKFusion::set_largest(const int64_t largest) { (void)this->AddAttr(kLargest, MakeValue(largest)); } +void TopKFusion::set_largest(const int64_t largest) { (void)this->AddAttr(kLargest, api::MakeValue(largest)); } int64_t TopKFusion::get_axis() const { auto value_ptr = GetAttr(kAxis); diff --git a/mindspore/core/ops/fusion/topk_fusion.h b/mindspore/core/ops/fusion/topk_fusion.h index 3e4994a1d0..6ff24f02ad 100644 --- a/mindspore/core/ops/fusion/topk_fusion.h +++ b/mindspore/core/ops/fusion/topk_fusion.h @@ -19,23 +19,18 @@ #include #include "ops/topk.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTopKFusion = "TopKFusion"; /// \brief TopKFusion defined TopK operator prototype of lite. -class MS_CORE_API TopKFusion : public TopK { +class MIND_API TopKFusion : public TopK { public: + MIND_API_BASE_MEMBER(TopKFusion); /// \brief Constructor. TopKFusion() : TopK(kNameTopKFusion) {} - /// \brief Destructor. - ~TopKFusion() = default; - - MS_DECLARE_PARENT(TopKFusion, TopK); - /// \brief Method to init the op's attributes. /// /// \param[in] sorted Define a boolean value indicate whether the output should be sorted. diff --git a/mindspore/core/ops/gather.cc b/mindspore/core/ops/gather.cc index 348b6c45e7..0abe3e57f2 100644 --- a/mindspore/core/ops/gather.cc +++ b/mindspore/core/ops/gather.cc @@ -21,10 +21,13 @@ #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" +#include "utils/check_convert_utils.h" namespace mindspore { namespace ops { // gather +MIND_API_BASE_IMPL(Gather, PrimitiveC, BaseOperator); AbstractBasePtr GatherInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/gather.h b/mindspore/core/ops/gather.h index 84d50eb755..f0be54d37f 100644 --- a/mindspore/core/ops/gather.h +++ b/mindspore/core/ops/gather.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGather = "Gather"; /// \brief Returns a slice of the input tensor based on the specified indices and axis. /// Refer to Python API @ref mindspore.ops.Gather for more details. -class MS_CORE_API Gather : public PrimitiveC { +class MIND_API Gather : public BaseOperator { public: + MIND_API_BASE_MEMBER(Gather); /// \brief Constructor. - Gather() : PrimitiveC(kNameGather) { InitIOName({"param", "indices", "axis"}, {"output"}); } - /// \brief Destructor. - ~Gather() = default; - MS_DECLARE_PARENT(Gather, PrimitiveC); + Gather() : BaseOperator(kNameGather) { InitIOName({"param", "indices", "axis"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Gather for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/gather_d.cc b/mindspore/core/ops/gather_d.cc index fd8db2ce93..1ee32a979c 100644 --- a/mindspore/core/ops/gather_d.cc +++ b/mindspore/core/ops/gather_d.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -59,6 +60,8 @@ TypePtr GatherDInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/gather_d.h b/mindspore/core/ops/gather_d.h index 834f5ee414..62a8cd7c7e 100644 --- a/mindspore/core/ops/gather_d.h +++ b/mindspore/core/ops/gather_d.h @@ -20,22 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Gathers values along an axis specified by dimension. /// Refer to Python API @ref mindspore.ops.GatherD for more details. -class MS_CORE_API GatherD : public PrimitiveC { +class MIND_API GatherD : public BaseOperator { public: + MIND_API_BASE_MEMBER(GatherD); /// \brief Constructor. - GatherD() : PrimitiveC(prim::kPrimGatherD->name()) { InitIOName({"x", "dim", "index"}, {"output"}); } - /// \brief Destructor. - ~GatherD() = default; - MS_DECLARE_PARENT(GatherD, PrimitiveC); + GatherD() : BaseOperator("GatherD") { InitIOName({"x", "dim", "index"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.GatherD for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/gather_nd.cc b/mindspore/core/ops/gather_nd.cc index 48446c9ec8..3f8dee5134 100644 --- a/mindspore/core/ops/gather_nd.cc +++ b/mindspore/core/ops/gather_nd.cc @@ -23,6 +23,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -83,6 +84,8 @@ TypePtr GatherNdInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/gather_nd.h b/mindspore/core/ops/gather_nd.h index d9a0522d58..02f0c3094e 100644 --- a/mindspore/core/ops/gather_nd.h +++ b/mindspore/core/ops/gather_nd.h @@ -18,26 +18,23 @@ #define MINDSPORE_CORE_OPS_GATHER_ND_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGatherNd = "GatherNd"; /// \brief Gathers slices from a tensor by indices. Refer to Python API @ref mindspore.ops.GatherNd for more details. -class MS_CORE_API GatherNd : public PrimitiveC { +class MIND_API GatherNd : public BaseOperator { public: + MIND_API_BASE_MEMBER(GatherNd); /// \brief Constructor. - GatherNd() : PrimitiveC(kNameGatherNd) { InitIOName({"x1", "x2"}, {"y"}); } - /// \brief Destructor. - ~GatherNd() = default; - MS_DECLARE_PARENT(GatherNd, PrimitiveC); + GatherNd() : BaseOperator(kNameGatherNd) { InitIOName({"x1", "x2"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.GatherNd for the inputs. void Init() const {} }; -AbstractBasePtr GatherNdInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GatherNdInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimGatherNdPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/gelu.cc b/mindspore/core/ops/gelu.cc index 0dd0188b37..4d37d88b86 100644 --- a/mindspore/core/ops/gelu.cc +++ b/mindspore/core/ops/gelu.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,6 +45,8 @@ TypePtr GeLUInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(GeLU, PrimitiveC, BaseOperator); AbstractBasePtr GeLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/gelu.h b/mindspore/core/ops/gelu.h index 2b6744eda4..edff2eb23c 100644 --- a/mindspore/core/ops/gelu.h +++ b/mindspore/core/ops/gelu.h @@ -17,23 +17,19 @@ #define MINDSPORE_CORE_OPS_GELU_H_ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameGeLU = prim::kGeLU; +constexpr auto kNameGeLU = "GeLU"; /// \brief Gaussian Error Linear Units activation function. /// Refer to Python API @ref mindspore.ops.GeLU for more details. -class MS_CORE_API GeLU : public PrimitiveC { +class MIND_API GeLU : public BaseOperator { public: + MIND_API_BASE_MEMBER(GeLU); /// \brief Constructor. - GeLU() : PrimitiveC(kNameGeLU) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~GeLU() = default; - MS_DECLARE_PARENT(GeLU, PrimitiveC); + GeLU() : BaseOperator(kNameGeLU) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.GeLU for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/ger.cc b/mindspore/core/ops/ger.cc index 986173dd0f..4ee7a7f6b6 100644 --- a/mindspore/core/ops/ger.cc +++ b/mindspore/core/ops/ger.cc @@ -20,6 +20,8 @@ #include "ops/ger.h" #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -53,6 +55,8 @@ TypePtr GerInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/ger.h b/mindspore/core/ops/ger.h index 43d61056aa..5f5f614711 100644 --- a/mindspore/core/ops/ger.h +++ b/mindspore/core/ops/ger.h @@ -21,26 +21,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGer = "Ger"; /// \brief Ger product of `x1` and `x2`. Calculate the outer product of two one-dimensional arrays. /// Refer to Python API @ref mindspore.ops.Ger for more details. -class MS_CORE_API Ger : public PrimitiveC { +class MIND_API Ger : public BaseOperator { public: + MIND_API_BASE_MEMBER(Ger); /// \brief Constructor. - Ger() : PrimitiveC(kNameGer) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Ger() = default; - MS_DECLARE_PARENT(Ger, PrimitiveC); + Ger() : BaseOperator(kNameGer) { InitIOName({"x", "y"}, {"output"}); } }; -AbstractBasePtr GerInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GerInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimGerPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/getnext.cc b/mindspore/core/ops/getnext.cc index a59943f986..b7077cf036 100644 --- a/mindspore/core/ops/getnext.cc +++ b/mindspore/core/ops/getnext.cc @@ -25,6 +25,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -86,6 +87,7 @@ abstract::AbstractBasePtr GetnextInferShape(const PrimitivePtr &primitive) { } } // namespace +MIND_API_BASE_IMPL(GetNext, PrimitiveC, BaseOperator); AbstractBasePtr GetNextInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return GetnextInferShape(primitive); diff --git a/mindspore/core/ops/getnext.h b/mindspore/core/ops/getnext.h index f83bf571a6..71a5371abd 100644 --- a/mindspore/core/ops/getnext.h +++ b/mindspore/core/ops/getnext.h @@ -21,27 +21,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameGetNext = prim::kGetNext; +constexpr auto kNameGetNext = "GetNext"; /// \brief Returns the next element in the dataset queue. /// Refer to Python API @ref mindspore.ops.GetNext for more details. -class MS_CORE_API GetNext : public PrimitiveC { +class MIND_API GetNext : public BaseOperator { public: + MIND_API_BASE_MEMBER(GetNext); /// \brief Constructor. - GetNext() : PrimitiveC(prim::kPrimGetNext->name()) {} - /// \brief Destructor. - ~GetNext() = default; - MS_DECLARE_PARENT(GetNext, PrimitiveC); + GetNext() : BaseOperator("GetNext") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.GetNext for the inputs. void Init() const {} }; -AbstractBasePtr GetNextInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GetNextInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_GETNEXT_H_ diff --git a/mindspore/core/ops/glu.cc b/mindspore/core/ops/glu.cc index 4b3679c0a1..c4f1faa246 100644 --- a/mindspore/core/ops/glu.cc +++ b/mindspore/core/ops/glu.cc @@ -15,12 +15,15 @@ */ #include "ops/glu.h" #include "ir/dtype/tensor_type.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(GLU, PrimitiveC, BaseOperator); void GLU::Init(int64_t axis) { set_axis(axis); } -void GLU::set_axis(int64_t axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void GLU::set_axis(int64_t axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } int64_t GLU::get_axis() const { auto value_ptr = GetAttr(kAxis); diff --git a/mindspore/core/ops/glu.h b/mindspore/core/ops/glu.h index b47735f012..ef780c7cca 100644 --- a/mindspore/core/ops/glu.h +++ b/mindspore/core/ops/glu.h @@ -18,19 +18,16 @@ #define MINDSPORE_CORE_OPS_GLU_H_ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameGLU = prim::kGLU; -class MS_CORE_API GLU : public PrimitiveC { +constexpr auto kNameGLU = "GLU"; +class MIND_API GLU : public BaseOperator { public: - GLU() : PrimitiveC(kNameGLU) { InitIOName({"x"}, {"output"}); } - ~GLU() = default; - MS_DECLARE_PARENT(GLU, PrimitiveC); + MIND_API_BASE_MEMBER(GLU); + GLU() : BaseOperator(kNameGLU) { InitIOName({"x"}, {"output"}); } void Init(int64_t axis); void set_axis(int64_t axis); int64_t get_axis() const; diff --git a/mindspore/core/ops/grad/abs_grad.cc b/mindspore/core/ops/grad/abs_grad.cc index 672087343a..d622944ed6 100644 --- a/mindspore/core/ops/grad/abs_grad.cc +++ b/mindspore/core/ops/grad/abs_grad.cc @@ -20,6 +20,8 @@ #include "abstract/param_validator.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -55,6 +57,7 @@ TypePtr AbsGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = AbsGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/abs_grad.h b/mindspore/core/ops/grad/abs_grad.h index 440d62f8ec..8042390fc8 100644 --- a/mindspore/core/ops/grad/abs_grad.h +++ b/mindspore/core/ops/grad/abs_grad.h @@ -20,19 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAbsGrad = "AbsGrad"; -class MS_CORE_API AbsGrad : public PrimitiveC { +class MIND_API AbsGrad : public BaseOperator { public: - AbsGrad() : PrimitiveC(kNameAbsGrad) {} - ~AbsGrad() = default; - MS_DECLARE_PARENT(AbsGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AbsGrad); + AbsGrad() : BaseOperator(kNameAbsGrad) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/acos_grad.cc b/mindspore/core/ops/grad/acos_grad.cc index 4c498be500..71a209bf89 100644 --- a/mindspore/core/ops/grad/acos_grad.cc +++ b/mindspore/core/ops/grad/acos_grad.cc @@ -15,6 +15,13 @@ */ #include "ops/grad/acos_grad.h" +#include +#include +#include "abstract/param_validator.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" +#include "ops/op_utils.h" namespace mindspore { namespace ops { @@ -40,6 +47,7 @@ TypePtr ACosGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/acos_grad.h b/mindspore/core/ops/grad/acos_grad.h index 44938646a6..5df2a851e7 100644 --- a/mindspore/core/ops/grad/acos_grad.h +++ b/mindspore/core/ops/grad/acos_grad.h @@ -22,25 +22,21 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameACosGrad = "ACosGrad"; -class ACosGrad : public PrimitiveC { +class MIND_API ACosGrad : public BaseOperator { public: - ACosGrad() : PrimitiveC(kNameACosGrad) { InitIOName({"y", "dy"}, {"z"}); } - ~ACosGrad() = default; - - MS_DECLARE_PARENT(ACosGrad, PrimitiveC); + MIND_API_BASE_MEMBER(ACosGrad); + ACosGrad() : BaseOperator(kNameACosGrad) { InitIOName({"y", "dy"}, {"z"}); } }; -AbstractBasePtr ACosGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ACosGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimACosGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/acosh_grad.cc b/mindspore/core/ops/grad/acosh_grad.cc index 9a2fd5a26d..c5d9ff07c2 100644 --- a/mindspore/core/ops/grad/acosh_grad.cc +++ b/mindspore/core/ops/grad/acosh_grad.cc @@ -15,6 +15,13 @@ */ #include "ops/grad/acosh_grad.h" +#include +#include +#include "ops/op_utils.h" +#include "abstract/param_validator.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +47,7 @@ TypePtr AcoshGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/acosh_grad.h b/mindspore/core/ops/grad/acosh_grad.h index 41e4b93773..be5c88efa6 100644 --- a/mindspore/core/ops/grad/acosh_grad.h +++ b/mindspore/core/ops/grad/acosh_grad.h @@ -22,25 +22,21 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAcoshGrad = "AcoshGrad"; -class AcoshGrad : public PrimitiveC { +class MIND_API AcoshGrad : public BaseOperator { public: - AcoshGrad() : PrimitiveC(kNameAcoshGrad) { InitIOName({"y", "dy"}, {"z"}); } - ~AcoshGrad() = default; - - MS_DECLARE_PARENT(AcoshGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AcoshGrad); + AcoshGrad() : BaseOperator(kNameAcoshGrad) { InitIOName({"y", "dy"}, {"z"}); } }; -AbstractBasePtr AcoshGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AcoshGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAcoshGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/activation_grad.cc b/mindspore/core/ops/grad/activation_grad.cc index 3a085d8c69..ee746d3a6e 100644 --- a/mindspore/core/ops/grad/activation_grad.cc +++ b/mindspore/core/ops/grad/activation_grad.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ActivationGrad, PrimitiveC, BaseOperator); void ActivationGrad::Init(const ActivationType &type, const float alpha) { this->set_activation_type(type); this->set_alpha(alpha); @@ -33,7 +35,7 @@ void ActivationGrad::Init(const ActivationType &type, const float alpha) { void ActivationGrad::set_activation_type(const ActivationType &type) { int64_t swi = type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType ActivationGrad::get_activation_type() const { @@ -42,7 +44,7 @@ ActivationType ActivationGrad::get_activation_type() const { return ActivationType(GetValue(value_ptr)); } -void ActivationGrad::set_alpha(const float alpha) { (void)this->AddAttr(kAlpha, MakeValue(alpha)); } +void ActivationGrad::set_alpha(const float alpha) { (void)this->AddAttr(kAlpha, api::MakeValue(alpha)); } float ActivationGrad::get_alpha() const { auto value_ptr = GetAttr(kAlpha); diff --git a/mindspore/core/ops/grad/activation_grad.h b/mindspore/core/ops/grad/activation_grad.h index 9ef3709198..faa7f27a5e 100644 --- a/mindspore/core/ops/grad/activation_grad.h +++ b/mindspore/core/ops/grad/activation_grad.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameActivationGrad = "ActivationGrad"; -class MS_CORE_API ActivationGrad : public PrimitiveC { +class MIND_API ActivationGrad : public BaseOperator { public: - ActivationGrad() : PrimitiveC(kNameActivationGrad) {} - ~ActivationGrad() = default; - MS_DECLARE_PARENT(ActivationGrad, PrimitiveC); + MIND_API_BASE_MEMBER(ActivationGrad); + ActivationGrad() : BaseOperator(kNameActivationGrad) {} void Init(const ActivationType &type = NO_ACTIVATION, const float alpha = 0.2); void set_activation_type(const ActivationType &type); void set_alpha(const float alpha); diff --git a/mindspore/core/ops/grad/add_grad.cc b/mindspore/core/ops/grad/add_grad.cc index c6a73cd179..342df58902 100644 --- a/mindspore/core/ops/grad/add_grad.cc +++ b/mindspore/core/ops/grad/add_grad.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(AddGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameAddGrad, AddGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/add_grad.h b/mindspore/core/ops/grad/add_grad.h index aa1b25440d..552b3431e8 100644 --- a/mindspore/core/ops/grad/add_grad.h +++ b/mindspore/core/ops/grad/add_grad.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAddGrad = "AddGrad"; -class MS_CORE_API AddGrad : public PrimitiveC { +class MIND_API AddGrad : public BaseOperator { public: - AddGrad() : PrimitiveC(kNameAddGrad) {} - ~AddGrad() = default; - MS_DECLARE_PARENT(AddGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AddGrad); + AddGrad() : BaseOperator(kNameAddGrad) {} void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/asin_grad.cc b/mindspore/core/ops/grad/asin_grad.cc index 8c5c539bcc..5fedf99fec 100644 --- a/mindspore/core/ops/grad/asin_grad.cc +++ b/mindspore/core/ops/grad/asin_grad.cc @@ -15,6 +15,11 @@ */ #include "ops/grad/asin_grad.h" +#include "abstract/param_validator.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +45,7 @@ TypePtr AsinGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/asin_grad.h b/mindspore/core/ops/grad/asin_grad.h index 1feb1ee497..a413894f6a 100644 --- a/mindspore/core/ops/grad/asin_grad.h +++ b/mindspore/core/ops/grad/asin_grad.h @@ -22,25 +22,22 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAsinGrad = "AsinGrad"; -class AsinGrad : public PrimitiveC { +class MIND_API AsinGrad : public BaseOperator { public: - AsinGrad() : PrimitiveC(kNameAsinGrad) { InitIOName({"y", "dy"}, {"z"}); } - ~AsinGrad() = default; - - MS_DECLARE_PARENT(AsinGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AsinGrad); + AsinGrad() : BaseOperator(kNameAsinGrad) { InitIOName({"y", "dy"}, {"z"}); } }; -AbstractBasePtr AsinGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AsinGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAsinGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/asinh_grad.cc b/mindspore/core/ops/grad/asinh_grad.cc index 08be8dd33e..447955c6b7 100644 --- a/mindspore/core/ops/grad/asinh_grad.cc +++ b/mindspore/core/ops/grad/asinh_grad.cc @@ -15,6 +15,9 @@ */ #include "ops/grad/asinh_grad.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +43,7 @@ TypePtr AsinhGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/asinh_grad.h b/mindspore/core/ops/grad/asinh_grad.h index a8b1bd10dc..4839d2d826 100644 --- a/mindspore/core/ops/grad/asinh_grad.h +++ b/mindspore/core/ops/grad/asinh_grad.h @@ -22,25 +22,21 @@ #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAsinhGrad = "AsinhGrad"; -class AsinhGrad : public PrimitiveC { +class MIND_API AsinhGrad : public BaseOperator { public: - AsinhGrad() : PrimitiveC(kNameAsinhGrad) { InitIOName({"y", "dy"}, {"z"}); } - ~AsinhGrad() = default; - - MS_DECLARE_PARENT(AsinhGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AsinhGrad); + AsinhGrad() : BaseOperator(kNameAsinhGrad) { InitIOName({"y", "dy"}, {"z"}); } }; -AbstractBasePtr AsinhGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AsinhGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimAsinhGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/atan_grad.cc b/mindspore/core/ops/grad/atan_grad.cc index 7fb2a075cb..2c310ba200 100644 --- a/mindspore/core/ops/grad/atan_grad.cc +++ b/mindspore/core/ops/grad/atan_grad.cc @@ -20,6 +20,7 @@ #include "abstract/param_validator.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -53,6 +54,7 @@ TypePtr AtanGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { auto type = AtanGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/atan_grad.h b/mindspore/core/ops/grad/atan_grad.h index 78afcd9afe..ec851822d3 100644 --- a/mindspore/core/ops/grad/atan_grad.h +++ b/mindspore/core/ops/grad/atan_grad.h @@ -20,19 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAtanGrad = "AtanGrad"; -class MS_CORE_API AtanGrad : public PrimitiveC { +class MIND_API AtanGrad : public BaseOperator { public: - AtanGrad() : PrimitiveC(kNameAtanGrad) {} - ~AtanGrad() = default; - MS_DECLARE_PARENT(AtanGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AtanGrad); + AtanGrad() : BaseOperator(kNameAtanGrad) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/avg_pool_3d_grad.cc b/mindspore/core/ops/grad/avg_pool_3d_grad.cc index bc7ac9eb08..1e7e99d24b 100644 --- a/mindspore/core/ops/grad/avg_pool_3d_grad.cc +++ b/mindspore/core/ops/grad/avg_pool_3d_grad.cc @@ -19,6 +19,9 @@ #include #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -58,6 +61,7 @@ TypePtr AvgPool3DGradInferType(const PrimitivePtr &primitive, const std::vector< } } // namespace +MIND_API_BASE_IMPL(AvgPool3DGrad, PrimitiveC, BaseOperator); AbstractBasePtr AvgPool3DGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto res = std::make_shared(AvgPool3DGradInferType(primitive, input_args), diff --git a/mindspore/core/ops/grad/avg_pool_3d_grad.h b/mindspore/core/ops/grad/avg_pool_3d_grad.h index 491ec77084..00f3b3967f 100644 --- a/mindspore/core/ops/grad/avg_pool_3d_grad.h +++ b/mindspore/core/ops/grad/avg_pool_3d_grad.h @@ -21,23 +21,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API AvgPool3DGrad : public PrimitiveC { +class MIND_API AvgPool3DGrad : public BaseOperator { public: - AvgPool3DGrad() : PrimitiveC(prim::kPrimAvgPool3DGrad->name()) { - InitIOName({"origin_input_size", "grad"}, {"output"}); - } - ~AvgPool3DGrad() = default; - MS_DECLARE_PARENT(AvgPool3DGrad, PrimitiveC); + MIND_API_BASE_MEMBER(AvgPool3DGrad); + AvgPool3DGrad() : BaseOperator("AvgPool3DGrad") { InitIOName({"origin_input_size", "grad"}, {"output"}); } }; -AbstractBasePtr AvgPool3DGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AvgPool3DGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/avg_pool_grad.cc b/mindspore/core/ops/grad/avg_pool_grad.cc index 9f02152bb5..2c00019ee1 100644 --- a/mindspore/core/ops/grad/avg_pool_grad.cc +++ b/mindspore/core/ops/grad/avg_pool_grad.cc @@ -16,9 +16,11 @@ #include "ops/grad/avg_pool_grad.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(AvgPoolGrad, PrimitiveC, PoolGrad); REGISTER_PRIMITIVE_C(kNameAvgPoolGrad, AvgPoolGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/avg_pool_grad.h b/mindspore/core/ops/grad/avg_pool_grad.h index d90a316cc1..e068067132 100644 --- a/mindspore/core/ops/grad/avg_pool_grad.h +++ b/mindspore/core/ops/grad/avg_pool_grad.h @@ -21,22 +21,20 @@ #include #include #include "ops/grad/pool_grad.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameAvgPoolGrad = "AvgPoolGrad"; -class MS_CORE_API AvgPoolGrad : public PoolGrad { +class MIND_API AvgPoolGrad : public PoolGrad { public: + MIND_API_BASE_MEMBER(AvgPoolGrad); AvgPoolGrad() : PoolGrad(kNameAvgPoolGrad) { InitIOName({"x_origin", "out_origin", "grad"}, {"output"}); } - ~AvgPoolGrad() = default; - MS_DECLARE_PARENT(AvgPoolGrad, PoolGrad); }; -AbstractBasePtr AvgPoolGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr AvgPoolGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/batch_norm_grad.cc b/mindspore/core/ops/grad/batch_norm_grad.cc index 7b558f3eb5..475b9b1909 100644 --- a/mindspore/core/ops/grad/batch_norm_grad.cc +++ b/mindspore/core/ops/grad/batch_norm_grad.cc @@ -19,15 +19,17 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(BatchNormGrad, PrimitiveC, BaseOperator); void BatchNormGrad::Init(const bool is_training, const float epsilon) { this->set_is_training(is_training); this->set_epsilon(epsilon); } -void BatchNormGrad::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, MakeValue(epsilon)); } +void BatchNormGrad::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, api::MakeValue(epsilon)); } float BatchNormGrad::get_epsilon() const { auto value_ptr = this->GetAttr(kEpsilon); @@ -36,7 +38,7 @@ float BatchNormGrad::get_epsilon() const { } void BatchNormGrad::set_is_training(const bool is_training) { - (void)this->AddAttr(kIsTraining, MakeValue(is_training)); + (void)this->AddAttr(kIsTraining, api::MakeValue(is_training)); } bool BatchNormGrad::get_is_training() const { diff --git a/mindspore/core/ops/grad/batch_norm_grad.h b/mindspore/core/ops/grad/batch_norm_grad.h index 4b34f1103e..53286fa263 100644 --- a/mindspore/core/ops/grad/batch_norm_grad.h +++ b/mindspore/core/ops/grad/batch_norm_grad.h @@ -18,18 +18,16 @@ #define MINDSPORE_CORE_OPS_BATCH_NORM_GRAD_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBatchNormGrad = "BatchNormGrad"; -class MS_CORE_API BatchNormGrad : public PrimitiveC { +class MIND_API BatchNormGrad : public BaseOperator { public: - BatchNormGrad() : PrimitiveC(kNameBatchNormGrad) {} - ~BatchNormGrad() = default; - MS_DECLARE_PARENT(BatchNormGrad, PrimitiveC); + MIND_API_BASE_MEMBER(BatchNormGrad); + BatchNormGrad() : BaseOperator(kNameBatchNormGrad) {} void Init(const bool is_training = false, const float epsilon = 1e-05); void set_is_training(const bool is_training); void set_epsilon(const float epsilon); @@ -37,8 +35,8 @@ class MS_CORE_API BatchNormGrad : public PrimitiveC { float get_epsilon() const; }; -AbstractBasePtr BatchNormGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BatchNormGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/bias_add_grad.cc b/mindspore/core/ops/grad/bias_add_grad.cc index 07804921b0..2072f1589e 100644 --- a/mindspore/core/ops/grad/bias_add_grad.cc +++ b/mindspore/core/ops/grad/bias_add_grad.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -71,6 +72,8 @@ TypePtr BiasAddGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(BiasAddGradInferShape(primitive, input_args), diff --git a/mindspore/core/ops/grad/bias_add_grad.h b/mindspore/core/ops/grad/bias_add_grad.h index 0bc4e7cc50..0574a05142 100644 --- a/mindspore/core/ops/grad/bias_add_grad.h +++ b/mindspore/core/ops/grad/bias_add_grad.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameBiasAddGrad = prim::kBiasAddGrad; -class MS_CORE_API BiasAddGrad : public PrimitiveC { +constexpr auto kNameBiasAddGrad = "BiasAddGrad"; +class MIND_API BiasAddGrad : public BaseOperator { public: - BiasAddGrad() : PrimitiveC(prim::kPrimBiasAddGrad->name()) { InitIOName({"x"}, {"output"}); } - ~BiasAddGrad() = default; - MS_DECLARE_PARENT(BiasAddGrad, PrimitiveC); + MIND_API_BASE_MEMBER(BiasAddGrad); + BiasAddGrad() : BaseOperator(kNameBiasAddGrad) { InitIOName({"x"}, {"output"}); } void Init() const {} }; -AbstractBasePtr BiasAddGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BiasAddGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/binary_cross_entropy_grad.cc b/mindspore/core/ops/grad/binary_cross_entropy_grad.cc index 36a38c7fbe..48a90397c4 100644 --- a/mindspore/core/ops/grad/binary_cross_entropy_grad.cc +++ b/mindspore/core/ops/grad/binary_cross_entropy_grad.cc @@ -19,6 +19,9 @@ #include #include "ops/grad/binary_cross_entropy_grad.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -26,13 +29,14 @@ void BinaryCrossEntropyGrad::Init(const Reduction &reduction) { set_reduction(re void BinaryCrossEntropyGrad::set_reduction(const Reduction &reduction) { int64_t swi = reduction; - (void)this->AddAttr(kReduction, MakeValue(swi)); + (void)this->AddAttr(kReduction, api::MakeValue(swi)); } Reduction BinaryCrossEntropyGrad::get_reduction() const { auto value_ptr = GetAttr(kReduction); return Reduction(GetValue(value_ptr)); } +MIND_API_BASE_IMPL(BinaryCrossEntropyGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameBinaryCrossEntropyGrad, BinaryCrossEntropyGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/binary_cross_entropy_grad.h b/mindspore/core/ops/grad/binary_cross_entropy_grad.h index d975e0d5af..0337846def 100644 --- a/mindspore/core/ops/grad/binary_cross_entropy_grad.h +++ b/mindspore/core/ops/grad/binary_cross_entropy_grad.h @@ -19,25 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBinaryCrossEntropyGrad = "BinaryCrossEntropyGrad"; -class MS_CORE_API BinaryCrossEntropyGrad : public PrimitiveC { +class MIND_API BinaryCrossEntropyGrad : public BaseOperator { public: - BinaryCrossEntropyGrad() : PrimitiveC(kNameBinaryCrossEntropyGrad) {} - ~BinaryCrossEntropyGrad() = default; - MS_DECLARE_PARENT(BinaryCrossEntropyGrad, PrimitiveC); + MIND_API_BASE_MEMBER(BinaryCrossEntropyGrad); + BinaryCrossEntropyGrad() : BaseOperator(kNameBinaryCrossEntropyGrad) {} void Init(const Reduction &reduction = MEAN); void set_reduction(const Reduction &reduction); Reduction get_reduction() const; }; -AbstractBasePtr BinaryCrossEntropyGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BinaryCrossEntropyGradInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_BINARY_CROSS_ENTROPY_GRAD_H_ diff --git a/mindspore/core/ops/grad/bn_grad.cc b/mindspore/core/ops/grad/bn_grad.cc index 2ec44e975d..2c3d0a7525 100644 --- a/mindspore/core/ops/grad/bn_grad.cc +++ b/mindspore/core/ops/grad/bn_grad.cc @@ -18,15 +18,17 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(BNGrad, PrimitiveC, BaseOperator); void BNGrad::Init(const float eps, const float momentum) { this->set_eps(eps); this->set_momentum(momentum); } -void BNGrad::set_eps(const float eps) { (void)this->AddAttr(kEps, MakeValue(eps)); } +void BNGrad::set_eps(const float eps) { (void)this->AddAttr(kEps, api::MakeValue(eps)); } float BNGrad::get_eps() const { auto value_ptr = this->GetAttr(kEps); @@ -34,7 +36,7 @@ float BNGrad::get_eps() const { return GetValue(value_ptr); } -void BNGrad::set_momentum(const float momentum) { (void)this->AddAttr(kMomentum, MakeValue(momentum)); } +void BNGrad::set_momentum(const float momentum) { (void)this->AddAttr(kMomentum, api::MakeValue(momentum)); } float BNGrad::get_momentum() const { auto value_ptr = this->GetAttr(kMomentum); diff --git a/mindspore/core/ops/grad/bn_grad.h b/mindspore/core/ops/grad/bn_grad.h index 38ce31f6bd..f677a3dacb 100644 --- a/mindspore/core/ops/grad/bn_grad.h +++ b/mindspore/core/ops/grad/bn_grad.h @@ -17,18 +17,16 @@ #ifndef MINDSPORE_CORE_OPS_BN_GRAD_H_ #define MINDSPORE_CORE_OPS_BN_GRAD_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBNGrad = "BNGrad"; -class MS_CORE_API BNGrad : public PrimitiveC { +class MIND_API BNGrad : public BaseOperator { public: - BNGrad() : PrimitiveC(kNameBNGrad) {} - ~BNGrad() = default; - MS_DECLARE_PARENT(BNGrad, PrimitiveC); + MIND_API_BASE_MEMBER(BNGrad); + BNGrad() : BaseOperator(kNameBNGrad) {} void Init(const float eps, const float momentum); void set_eps(const float eps); void set_momentum(const float momentum); diff --git a/mindspore/core/ops/grad/bn_training_reduce_grad.cc b/mindspore/core/ops/grad/bn_training_reduce_grad.cc index a7c18375f6..34ab4f478d 100644 --- a/mindspore/core/ops/grad/bn_training_reduce_grad.cc +++ b/mindspore/core/ops/grad/bn_training_reduce_grad.cc @@ -17,6 +17,7 @@ #include "ops/grad/bn_training_reduce_grad.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,9 +45,10 @@ TypePtr BNTrainingReduceGradInferType(const PrimitivePtr &prim, const std::vecto } } // namespace +MIND_API_BASE_IMPL(BNTrainingReduceGrad, PrimitiveC, BaseOperator); void BNTrainingReduceGrad::Init(const float epsilon) { this->set_epsilon(epsilon); } -void BNTrainingReduceGrad::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, MakeValue(epsilon)); } +void BNTrainingReduceGrad::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, api::MakeValue(epsilon)); } float BNTrainingReduceGrad::get_epsilon() const { auto value_ptr = GetAttr(kEpsilon); diff --git a/mindspore/core/ops/grad/bn_training_reduce_grad.h b/mindspore/core/ops/grad/bn_training_reduce_grad.h index a377e2ed09..e22fa5111c 100644 --- a/mindspore/core/ops/grad/bn_training_reduce_grad.h +++ b/mindspore/core/ops/grad/bn_training_reduce_grad.h @@ -20,30 +20,28 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBNTrainingReduceGrad = "BNTrainingReduceGrad"; -class BNTrainingReduceGrad : public PrimitiveC { +class MIND_API BNTrainingReduceGrad : public BaseOperator { public: - BNTrainingReduceGrad() : PrimitiveC(kNameBNTrainingReduceGrad) { + MIND_API_BASE_MEMBER(BNTrainingReduceGrad); + BNTrainingReduceGrad() : BaseOperator(kNameBNTrainingReduceGrad) { InitIOName({"grads", "x", "diff_scale", "diff_offset", "scale", "batch_mean", "batch_variance"}, {"y"}); } - explicit BNTrainingReduceGrad(const std::string k_name) : PrimitiveC(k_name) { + explicit BNTrainingReduceGrad(const std::string k_name) : BaseOperator(k_name) { InitIOName({"grads", "x", "diff_scale", "diff_offset", "scale", "batch_mean", "batch_variance"}, {"y"}); } - ~BNTrainingReduceGrad() = default; - MS_DECLARE_PARENT(BNTrainingReduceGrad, PrimitiveC); void Init(const float epsilon = 0.0001); void set_epsilon(const float epsilon); float get_epsilon() const; }; -AbstractBasePtr BNTrainingReduceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BNTrainingReduceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/bn_training_update_grad.cc b/mindspore/core/ops/grad/bn_training_update_grad.cc index 3eab005715..653d42e53b 100644 --- a/mindspore/core/ops/grad/bn_training_update_grad.cc +++ b/mindspore/core/ops/grad/bn_training_update_grad.cc @@ -22,6 +22,8 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -49,6 +51,7 @@ TuplePtr BNTrainingUpdateGradInferType(const PrimitivePtr &primitive, const std: } } // namespace +MIND_API_BASE_IMPL(BNTrainingUpdateGrad, PrimitiveC, BaseOperator); AbstractBasePtr BNTrainingUpdateGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/bn_training_update_grad.h b/mindspore/core/ops/grad/bn_training_update_grad.h index 62a9743568..69e95d1408 100644 --- a/mindspore/core/ops/grad/bn_training_update_grad.h +++ b/mindspore/core/ops/grad/bn_training_update_grad.h @@ -22,24 +22,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameBNTrainingUpdateGrad = "BNTrainingUpdateGrad"; -class MS_CORE_API BNTrainingUpdateGrad : public PrimitiveC { +class MIND_API BNTrainingUpdateGrad : public BaseOperator { public: - BNTrainingUpdateGrad() : PrimitiveC(kNameBNTrainingUpdateGrad) { + MIND_API_BASE_MEMBER(BNTrainingUpdateGrad); + BNTrainingUpdateGrad() : BaseOperator(kNameBNTrainingUpdateGrad) { InitIOName({"grads", "x", "batch_mean", "batch_variance"}, {"diff_scale", "diff_offset"}); } - ~BNTrainingUpdateGrad() = default; - MS_DECLARE_PARENT(BNTrainingUpdateGrad, PrimitiveC); }; -AbstractBasePtr BNTrainingUpdateGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr BNTrainingUpdateGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimBNTrainingUpdateGradPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/grad/cdist_grad.cc b/mindspore/core/ops/grad/cdist_grad.cc index a9d7a416cf..62629207eb 100644 --- a/mindspore/core/ops/grad/cdist_grad.cc +++ b/mindspore/core/ops/grad/cdist_grad.cc @@ -19,6 +19,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -59,6 +60,7 @@ TypePtr CdistGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/cdist_grad.h b/mindspore/core/ops/grad/cdist_grad.h index c8819f10e5..19930dcde4 100644 --- a/mindspore/core/ops/grad/cdist_grad.h +++ b/mindspore/core/ops/grad/cdist_grad.h @@ -22,22 +22,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameCdistGrad = "CdistGrad"; -class MS_CORE_API CdistGrad : public PrimitiveC { +class MIND_API CdistGrad : public BaseOperator { public: - CdistGrad() : PrimitiveC(kNameCdistGrad) { InitIOName({"grad", "input_x", "input_y", "cdist"}, {"output"}); } - ~CdistGrad() = default; - MS_DECLARE_PARENT(CdistGrad, PrimitiveC); + MIND_API_BASE_MEMBER(CdistGrad); + CdistGrad() : BaseOperator(kNameCdistGrad) { InitIOName({"grad", "input_x", "input_y", "cdist"}, {"output"}); } }; -AbstractBasePtr CdistGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr CdistGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/conv2d_backprop_filter.cc b/mindspore/core/ops/grad/conv2d_backprop_filter.cc index 2cdb3dc1d1..46c4621261 100644 --- a/mindspore/core/ops/grad/conv2d_backprop_filter.cc +++ b/mindspore/core/ops/grad/conv2d_backprop_filter.cc @@ -20,6 +20,8 @@ #include "ops/grad/conv2d_backprop_filter.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -100,6 +102,7 @@ TypePtr Conv2DBackpropFilterInferType(const PrimitivePtr &prim, const std::vecto } } // namespace +MIND_API_BASE_IMPL(Conv2DBackpropFilter, PrimitiveC, BaseOperator); void Conv2DBackpropFilter::Init(const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &pad_list, const int64_t mode, const std::vector &stride, const std::vector &dilation, @@ -121,7 +124,7 @@ void Conv2DBackpropFilter::Init(const int64_t out_channel, const std::vectorAddAttr(kOutChannel, MakeValue(out_channel)); + (void)this->AddAttr(kOutChannel, api::MakeValue(out_channel)); } int64_t Conv2DBackpropFilter::get_out_channel() const { @@ -131,7 +134,7 @@ int64_t Conv2DBackpropFilter::get_out_channel() const { } void Conv2DBackpropFilter::set_kernel_size(const std::vector &kernel_size) { - (void)this->AddAttr(kKernelSize, MakeValue(kernel_size)); + (void)this->AddAttr(kKernelSize, api::MakeValue(kernel_size)); } std::vector Conv2DBackpropFilter::get_kernel_size() const { @@ -142,7 +145,7 @@ std::vector Conv2DBackpropFilter::get_kernel_size() const { void Conv2DBackpropFilter::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode Conv2DBackpropFilter::get_pad_mode() const { @@ -152,7 +155,7 @@ PadMode Conv2DBackpropFilter::get_pad_mode() const { } void Conv2DBackpropFilter::set_pad_list(const std::vector &pad_list) { - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } std::vector Conv2DBackpropFilter::get_pad_list() const { @@ -161,7 +164,7 @@ std::vector Conv2DBackpropFilter::get_pad_list() const { return GetValue>(value_ptr); } -void Conv2DBackpropFilter::set_mode(const int64_t mode) { (void)this->AddAttr(kMode, MakeValue(mode)); } +void Conv2DBackpropFilter::set_mode(const int64_t mode) { (void)this->AddAttr(kMode, api::MakeValue(mode)); } int64_t Conv2DBackpropFilter::get_mode() const { auto value_ptr = GetAttr(kMode); @@ -170,7 +173,7 @@ int64_t Conv2DBackpropFilter::get_mode() const { } void Conv2DBackpropFilter::set_stride(const std::vector &stride) { - (void)this->AddAttr(kStride, MakeValue(stride)); + (void)this->AddAttr(kStride, api::MakeValue(stride)); } std::vector Conv2DBackpropFilter::get_stride() const { @@ -180,7 +183,7 @@ std::vector Conv2DBackpropFilter::get_stride() const { } void Conv2DBackpropFilter::set_dilation(const std::vector &dilation) { - (void)this->AddAttr(kDilation, MakeValue(dilation)); + (void)this->AddAttr(kDilation, api::MakeValue(dilation)); } std::vector Conv2DBackpropFilter::get_dilation() const { @@ -189,7 +192,7 @@ std::vector Conv2DBackpropFilter::get_dilation() const { return GetValue>(value_ptr); } -void Conv2DBackpropFilter::set_group(const int64_t group) { (void)this->AddAttr(kGroup, MakeValue(group)); } +void Conv2DBackpropFilter::set_group(const int64_t group) { (void)this->AddAttr(kGroup, api::MakeValue(group)); } int64_t Conv2DBackpropFilter::get_group() const { auto value_ptr = GetAttr(kGroup); @@ -199,7 +202,7 @@ int64_t Conv2DBackpropFilter::get_group() const { void Conv2DBackpropFilter::set_format(const Format &format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } Format Conv2DBackpropFilter::get_format() const { diff --git a/mindspore/core/ops/grad/conv2d_backprop_filter.h b/mindspore/core/ops/grad/conv2d_backprop_filter.h index e2ff2ed6e6..93636f01b7 100644 --- a/mindspore/core/ops/grad/conv2d_backprop_filter.h +++ b/mindspore/core/ops/grad/conv2d_backprop_filter.h @@ -20,23 +20,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameConv2DBackpropFilter = "Conv2DBackpropFilter"; -class MS_CORE_API Conv2DBackpropFilter : public PrimitiveC { +class MIND_API Conv2DBackpropFilter : public BaseOperator { public: - Conv2DBackpropFilter() : PrimitiveC(kNameConv2DBackpropFilter) { + MIND_API_BASE_MEMBER(Conv2DBackpropFilter); + Conv2DBackpropFilter() : BaseOperator(kNameConv2DBackpropFilter) { InitIOName({"out_backprop", "input", "filter_sizes"}, {"output"}); } - explicit Conv2DBackpropFilter(const std::string k_name) : PrimitiveC(k_name) { + explicit Conv2DBackpropFilter(const std::string k_name) : BaseOperator(k_name) { InitIOName({"out_backprop", "input", "filter_sizes"}, {"output"}); } - ~Conv2DBackpropFilter() = default; - MS_DECLARE_PARENT(Conv2DBackpropFilter, PrimitiveC); void Init(const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode = VALID, const std::vector &pad_list = {0, 0, 0, 0}, const int64_t mode = 1, const std::vector &stride = {1, 1}, const std::vector &dilation = {1, 1, 1, 1}, @@ -64,8 +63,8 @@ class MS_CORE_API Conv2DBackpropFilter : public PrimitiveC { int64_t get_group() const; Format get_format() const; }; -AbstractBasePtr Conv2DBackpropFilterInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr Conv2DBackpropFilterInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/conv2d_backprop_input.cc b/mindspore/core/ops/grad/conv2d_backprop_input.cc index 030cfb305c..5e9124ed3e 100644 --- a/mindspore/core/ops/grad/conv2d_backprop_input.cc +++ b/mindspore/core/ops/grad/conv2d_backprop_input.cc @@ -21,6 +21,9 @@ #include #include "ops/grad/conv2d_backprop_input.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -159,6 +162,8 @@ TypePtr Conv2DBackpropInputInferType(const PrimitivePtr &prim, const std::vector return CheckAndConvertUtils::CheckTensorTypeSame(types, valid_x_type, prim_name); } } // namespace + +MIND_API_BASE_IMPL(Conv2DBackpropInput, PrimitiveC, BaseOperator); AbstractBasePtr Conv2DBackpropInputInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); @@ -193,19 +198,20 @@ void Conv2DBackpropInput::Init(int64_t out_channel, const std::vector & void Conv2DBackpropInput::set_out_channel(int64_t out_channel) { (void)AddAttr(kOutChannel, - MakeValue(CheckAndConvertUtils::CheckInteger(kOutChannel, out_channel, kGreaterThan, 0, name()))); + api::MakeValue(CheckAndConvertUtils::CheckInteger(kOutChannel, out_channel, kGreaterThan, 0, name()))); } void Conv2DBackpropInput::set_kernel_size(const std::vector &kernel_size) { - (void)AddAttr(kKernelSize, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, name()))); + (void)AddAttr(kKernelSize, + api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, name()))); } void Conv2DBackpropInput::set_stride(const std::vector &stride) { - (void)AddAttr(kStride, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStride, stride, name()))); + (void)AddAttr(kStride, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStride, stride, name()))); } void Conv2DBackpropInput::set_dilation(const std::vector &dilation) { - (void)AddAttr(kDilation, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kDilation, dilation, name()))); + (void)AddAttr(kDilation, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kDilation, dilation, name()))); } void Conv2DBackpropInput::set_pad_mode(const PadMode &pad_mode) { @@ -218,30 +224,30 @@ void Conv2DBackpropInput::set_pad_mode(const PadMode &pad_mode) { CheckAndConvertUtils::Check(kPad, pad, kEqual, {0, 0, 0, 0}, name()); } int64_t swi = pad_mode; - (void)AddAttr(kPadMode, MakeValue(swi)); + (void)AddAttr(kPadMode, api::MakeValue(swi)); } void Conv2DBackpropInput::set_pad(const std::vector &pad) { const int64_t pad_size = 4; (void)CheckAndConvertUtils::CheckInteger("pad_size", SizeToLong(pad.size()), kEqual, pad_size, name()); - (void)AddAttr(kPad, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kPad, pad, name()))); + (void)AddAttr(kPad, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kPad, pad, name()))); } void Conv2DBackpropInput::set_mode(int64_t mode) { - (void)AddAttr(kMode, MakeValue(CheckAndConvertUtils::CheckInteger(kMode, mode, kEqual, 1, name()))); + (void)AddAttr(kMode, api::MakeValue(CheckAndConvertUtils::CheckInteger(kMode, mode, kEqual, 1, name()))); } void Conv2DBackpropInput::set_group(int64_t group) { - (void)AddAttr(kGroup, MakeValue(CheckAndConvertUtils::CheckInteger(kGroup, group, kGreaterThan, 0, name()))); + (void)AddAttr(kGroup, api::MakeValue(CheckAndConvertUtils::CheckInteger(kGroup, group, kGreaterThan, 0, name()))); } void Conv2DBackpropInput::set_format(const Format &format) { int64_t f = format; - (void)AddAttr(kFormat, MakeValue(f)); + (void)AddAttr(kFormat, api::MakeValue(f)); } void Conv2DBackpropInput::set_pad_list(const std::vector &pad_list) { - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } int64_t Conv2DBackpropInput::get_out_channel() const { diff --git a/mindspore/core/ops/grad/conv2d_backprop_input.h b/mindspore/core/ops/grad/conv2d_backprop_input.h index ef765eb095..0bc292b8c0 100644 --- a/mindspore/core/ops/grad/conv2d_backprop_input.h +++ b/mindspore/core/ops/grad/conv2d_backprop_input.h @@ -21,20 +21,19 @@ #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" + namespace mindspore { namespace ops { constexpr auto kNameConv2DBackpropInput = "Conv2DBackpropInput"; -class MS_CORE_API Conv2DBackpropInput : public PrimitiveC { +class MIND_API Conv2DBackpropInput : public BaseOperator { public: - explicit Conv2DBackpropInput(const std::string &k_name = kNameConv2DBackpropInput) : PrimitiveC(k_name) { + MIND_API_BASE_MEMBER(Conv2DBackpropInput); + explicit Conv2DBackpropInput(const std::string &k_name = kNameConv2DBackpropInput) : BaseOperator(k_name) { InitIOName({"out_backprop", "filter", "input_sizes"}, {"output"}); } - ~Conv2DBackpropInput() = default; - MS_DECLARE_PARENT(Conv2DBackpropInput, PrimitiveC); void Init(int64_t out_channel, const std::vector &kernel_size, int64_t mode = 1, const PadMode &pad_mode = VALID, const std::vector &pad = {0, 0, 0, 0}, const std::vector &stride = {1, 1, 1, 1}, const std::vector &dilation = {1, 1, 1, 1}, @@ -60,8 +59,8 @@ class MS_CORE_API Conv2DBackpropInput : public PrimitiveC { Format get_format() const; std::vector get_pad_list() const; }; -AbstractBasePtr Conv2DBackpropInputInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr Conv2DBackpropInputInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_CONV2D_BACKPROP_INPUT_H_ diff --git a/mindspore/core/ops/grad/de_conv2d_grad_filter.cc b/mindspore/core/ops/grad/de_conv2d_grad_filter.cc index 2d3f6e91eb..de09c49623 100644 --- a/mindspore/core/ops/grad/de_conv2d_grad_filter.cc +++ b/mindspore/core/ops/grad/de_conv2d_grad_filter.cc @@ -18,9 +18,11 @@ #include "ops/grad/de_conv2d_grad_filter.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(DeConv2DGradFilter, PrimitiveC, BaseOperator); void DeConv2DGradFilter::Init(const int64_t in_channel, const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &pad_list, const std::vector &stride, @@ -40,7 +42,7 @@ void DeConv2DGradFilter::Init(const int64_t in_channel, const int64_t out_channe } void DeConv2DGradFilter::set_in_channel(const int64_t in_channel) { - (void)this->AddAttr(kInChannel, MakeValue(in_channel)); + (void)this->AddAttr(kInChannel, api::MakeValue(in_channel)); } int64_t DeConv2DGradFilter::get_in_channel() const { @@ -50,7 +52,7 @@ int64_t DeConv2DGradFilter::get_in_channel() const { } void DeConv2DGradFilter::set_out_channel(const int64_t out_channel) { - (void)this->AddAttr(kOutChannel, MakeValue(out_channel)); + (void)this->AddAttr(kOutChannel, api::MakeValue(out_channel)); } int64_t DeConv2DGradFilter::get_out_channel() const { @@ -60,7 +62,7 @@ int64_t DeConv2DGradFilter::get_out_channel() const { } void DeConv2DGradFilter::set_kernel_size(const std::vector &kernel_size) { - (void)this->AddAttr(kKernelSize, MakeValue(kernel_size)); + (void)this->AddAttr(kKernelSize, api::MakeValue(kernel_size)); } std::vector DeConv2DGradFilter::get_kernel_size() const { @@ -71,7 +73,7 @@ std::vector DeConv2DGradFilter::get_kernel_size() const { void DeConv2DGradFilter::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode DeConv2DGradFilter::get_pad_mode() const { @@ -81,7 +83,7 @@ PadMode DeConv2DGradFilter::get_pad_mode() const { } void DeConv2DGradFilter::set_pad_list(const std::vector &pad_list) { - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } std::vector DeConv2DGradFilter::get_pad_list() const { @@ -91,7 +93,7 @@ std::vector DeConv2DGradFilter::get_pad_list() const { } void DeConv2DGradFilter::set_stride(const std::vector &stride) { - (void)this->AddAttr(kStride, MakeValue(stride)); + (void)this->AddAttr(kStride, api::MakeValue(stride)); } std::vector DeConv2DGradFilter::get_stride() const { @@ -101,7 +103,7 @@ std::vector DeConv2DGradFilter::get_stride() const { } void DeConv2DGradFilter::set_dilation(const std::vector &dilation) { - (void)this->AddAttr(kDilation, MakeValue(dilation)); + (void)this->AddAttr(kDilation, api::MakeValue(dilation)); } std::vector DeConv2DGradFilter::get_dilation() const { @@ -110,7 +112,7 @@ std::vector DeConv2DGradFilter::get_dilation() const { return GetValue>(value_ptr); } -void DeConv2DGradFilter::set_group(const int64_t group) { (void)this->AddAttr(kGroup, MakeValue(group)); } +void DeConv2DGradFilter::set_group(const int64_t group) { (void)this->AddAttr(kGroup, api::MakeValue(group)); } int64_t DeConv2DGradFilter::get_group() const { auto value_ptr = GetAttr(kGroup); @@ -120,7 +122,7 @@ int64_t DeConv2DGradFilter::get_group() const { void DeConv2DGradFilter::set_format(const Format &format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } Format DeConv2DGradFilter::get_format() const { @@ -131,7 +133,7 @@ Format DeConv2DGradFilter::get_format() const { void DeConv2DGradFilter::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType DeConv2DGradFilter::get_activation_type() const { @@ -140,7 +142,7 @@ ActivationType DeConv2DGradFilter::get_activation_type() const { return ActivationType(GetValue(value_ptr)); } -void DeConv2DGradFilter::set_has_bias(const bool has_bias) { (void)this->AddAttr(kHasBias, MakeValue(has_bias)); } +void DeConv2DGradFilter::set_has_bias(const bool has_bias) { (void)this->AddAttr(kHasBias, api::MakeValue(has_bias)); } bool DeConv2DGradFilter::get_has_bias() const { auto value_ptr = GetAttr(kHasBias); diff --git a/mindspore/core/ops/grad/de_conv2d_grad_filter.h b/mindspore/core/ops/grad/de_conv2d_grad_filter.h index 9c9be83281..e0a8bd2d2c 100644 --- a/mindspore/core/ops/grad/de_conv2d_grad_filter.h +++ b/mindspore/core/ops/grad/de_conv2d_grad_filter.h @@ -18,18 +18,17 @@ #define MINDSPORE_CORE_OPS_DE_CONV2D_GRAD_FILTER_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameDeConv2DGradFilter = "DeConv2DGradFilter"; -class MS_CORE_API DeConv2DGradFilter : public PrimitiveC { +class MIND_API DeConv2DGradFilter : public BaseOperator { public: - DeConv2DGradFilter() : PrimitiveC(kNameDeConv2DGradFilter) {} - ~DeConv2DGradFilter() = default; - MS_DECLARE_PARENT(DeConv2DGradFilter, PrimitiveC); + MIND_API_BASE_MEMBER(DeConv2DGradFilter); + DeConv2DGradFilter() : BaseOperator(kNameDeConv2DGradFilter) {} void Init(const int64_t in_channel, const int64_t out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &pad_list, const std::vector &stride, const std::vector &dilation, const int64_t group, const Format &format = NCHW, diff --git a/mindspore/core/ops/grad/div_grad.cc b/mindspore/core/ops/grad/div_grad.cc index 01a96f2312..8dc2b3911e 100644 --- a/mindspore/core/ops/grad/div_grad.cc +++ b/mindspore/core/ops/grad/div_grad.cc @@ -17,9 +17,11 @@ #include "ops/grad/div_grad.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(DivGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameDivGrad, DivGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/div_grad.h b/mindspore/core/ops/grad/div_grad.h index 4878c5f7c9..aaa58cd5a3 100644 --- a/mindspore/core/ops/grad/div_grad.h +++ b/mindspore/core/ops/grad/div_grad.h @@ -16,18 +16,16 @@ #ifndef MINDSPORE_CORE_OPS_DIV_GRAD_H_ #define MINDSPORE_CORE_OPS_DIV_GRAD_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameDivGrad = "DivGrad"; -class MS_CORE_API DivGrad : public PrimitiveC { +class MIND_API DivGrad : public BaseOperator { public: - DivGrad() : PrimitiveC(kNameDivGrad) {} - ~DivGrad() = default; - MS_DECLARE_PARENT(DivGrad, PrimitiveC); + MIND_API_BASE_MEMBER(DivGrad); + DivGrad() : BaseOperator(kNameDivGrad) {} void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/dropout_grad.cc b/mindspore/core/ops/grad/dropout_grad.cc index f21f6acfa2..39bb88f97e 100644 --- a/mindspore/core/ops/grad/dropout_grad.cc +++ b/mindspore/core/ops/grad/dropout_grad.cc @@ -21,14 +21,16 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(DropoutGrad, PrimitiveC, BaseOperator); void DropoutGrad::Init(const float keep_prob) { this->set_keep_prob(keep_prob); } void DropoutGrad::set_keep_prob(const float keep_prob) { CheckAndConvertUtils::CheckInRange(kKeepProb, keep_prob, kIncludeRight, {0.0, 1.0}, this->name()); - (void)this->AddAttr(kKeepProb, MakeValue(keep_prob)); + (void)this->AddAttr(kKeepProb, api::MakeValue(keep_prob)); } float DropoutGrad::get_keep_prob() const { diff --git a/mindspore/core/ops/grad/dropout_grad.h b/mindspore/core/ops/grad/dropout_grad.h index 57dcd2be22..87395415ef 100644 --- a/mindspore/core/ops/grad/dropout_grad.h +++ b/mindspore/core/ops/grad/dropout_grad.h @@ -19,24 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API DropoutGrad : public PrimitiveC { +class MIND_API DropoutGrad : public BaseOperator { public: - DropoutGrad() : PrimitiveC(prim::kPrimDropoutGrad->name()) {} - ~DropoutGrad() = default; - MS_DECLARE_PARENT(DropoutGrad, PrimitiveC); + MIND_API_BASE_MEMBER(DropoutGrad); + DropoutGrad() : BaseOperator("DropoutGrad") {} void Init(const float keep_prob = 0.5); void set_keep_prob(const float keep_prob); float get_keep_prob() const; }; -AbstractBasePtr DropoutGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr DropoutGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_DROPOUT_GRAD_H_ diff --git a/mindspore/core/ops/grad/einsum_grad.cc b/mindspore/core/ops/grad/einsum_grad.cc index a1c7cbbd98..620c5582bc 100644 --- a/mindspore/core/ops/grad/einsum_grad.cc +++ b/mindspore/core/ops/grad/einsum_grad.cc @@ -19,12 +19,14 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(EinsumGrad, PrimitiveC, BaseOperator); void EinsumGrad::Init(const std::string equation) { this->set_equation(equation); } -void EinsumGrad::set_equation(const std::string equation) { (void)this->AddAttr(kEquation, MakeValue(equation)); } +void EinsumGrad::set_equation(const std::string equation) { (void)this->AddAttr(kEquation, api::MakeValue(equation)); } std::string EinsumGrad::get_equation() const { auto value_ptr = this->GetAttr(kEquation); diff --git a/mindspore/core/ops/grad/einsum_grad.h b/mindspore/core/ops/grad/einsum_grad.h index 225536c8e9..739d79e225 100644 --- a/mindspore/core/ops/grad/einsum_grad.h +++ b/mindspore/core/ops/grad/einsum_grad.h @@ -19,18 +19,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameEinsumGrad = "EinsumGrad"; -class EinsumGrad : public PrimitiveC { +class MIND_API EinsumGrad : public BaseOperator { public: - EinsumGrad() : PrimitiveC(kNameEinsumGrad) {} - ~EinsumGrad() = default; - MS_DECLARE_PARENT(EinsumGrad, PrimitiveC); + MIND_API_BASE_MEMBER(EinsumGrad); + EinsumGrad() : BaseOperator(kNameEinsumGrad) {} void Init(const std::string equation); void set_equation(const std::string equation); std::string get_equation() const; diff --git a/mindspore/core/ops/grad/elu_grad.cc b/mindspore/core/ops/grad/elu_grad.cc index 2a6ae68b8f..04950f08c7 100644 --- a/mindspore/core/ops/grad/elu_grad.cc +++ b/mindspore/core/ops/grad/elu_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -63,6 +64,8 @@ TypePtr EluGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = EluGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/elu_grad.h b/mindspore/core/ops/grad/elu_grad.h index b12016e986..4218ed6a1c 100644 --- a/mindspore/core/ops/grad/elu_grad.h +++ b/mindspore/core/ops/grad/elu_grad.h @@ -19,19 +19,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameEluGrad = "EluGrad"; -class EluGrad : public PrimitiveC { +class MIND_API EluGrad : public BaseOperator { public: - EluGrad() : PrimitiveC(kNameEluGrad) {} - ~EluGrad() = default; - MS_DECLARE_PARENT(EluGrad, PrimitiveC); + MIND_API_BASE_MEMBER(EluGrad); + EluGrad() : BaseOperator(kNameEluGrad) {} void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/fast_gelu_grad.cc b/mindspore/core/ops/grad/fast_gelu_grad.cc index 81ee4e2125..0bdcc8d8ab 100644 --- a/mindspore/core/ops/grad/fast_gelu_grad.cc +++ b/mindspore/core/ops/grad/fast_gelu_grad.cc @@ -20,6 +20,7 @@ #include "abstract/param_validator.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -53,6 +54,7 @@ TypePtr FastGeLUGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = FastGeLUGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/fast_gelu_grad.h b/mindspore/core/ops/grad/fast_gelu_grad.h index 99c4cfdbc1..75590f5484 100644 --- a/mindspore/core/ops/grad/fast_gelu_grad.h +++ b/mindspore/core/ops/grad/fast_gelu_grad.h @@ -20,19 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFastGeLUGrad = "FastGeLUGrad"; -class FastGeLUGrad : public PrimitiveC { +class FastGeLUGrad : public BaseOperator { public: - FastGeLUGrad() : PrimitiveC(prim::kPrimFastGeLUGrad->name()) { InitIOName({"x"}, {"output"}); } - ~FastGeLUGrad() = default; - MS_DECLARE_PARENT(FastGeLUGrad, PrimitiveC); + MIND_API_BASE_MEMBER(FastGeLUGrad); + FastGeLUGrad() : BaseOperator("FastGeLUGrad") { InitIOName({"x"}, {"output"}); } }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/flatten_grad.cc b/mindspore/core/ops/grad/flatten_grad.cc index 016fc45a49..153d3e51a6 100644 --- a/mindspore/core/ops/grad/flatten_grad.cc +++ b/mindspore/core/ops/grad/flatten_grad.cc @@ -15,9 +15,13 @@ */ #include "ops/grad/flatten_grad.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(FlattenGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameFlattenGrad, FlattenGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/flatten_grad.h b/mindspore/core/ops/grad/flatten_grad.h index f492a53efb..be8595d973 100644 --- a/mindspore/core/ops/grad/flatten_grad.h +++ b/mindspore/core/ops/grad/flatten_grad.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameFlattenGrad = "FlattenGrad"; -class MS_CORE_API FlattenGrad : public PrimitiveC { +class MIND_API FlattenGrad : public BaseOperator { public: - FlattenGrad() : PrimitiveC(kNameFlattenGrad) { InitIOName({"x", "shape"}, {"output"}); } - ~FlattenGrad() = default; - MS_DECLARE_PARENT(FlattenGrad, PrimitiveC); + MIND_API_BASE_MEMBER(FlattenGrad); + FlattenGrad() : BaseOperator(kNameFlattenGrad) { InitIOName({"x", "shape"}, {"output"}); } }; -AbstractBasePtr FlattenGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr FlattenGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimFlattenGrad = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/gelu_grad.cc b/mindspore/core/ops/grad/gelu_grad.cc index a48f29535e..6e8e0b4b34 100644 --- a/mindspore/core/ops/grad/gelu_grad.cc +++ b/mindspore/core/ops/grad/gelu_grad.cc @@ -25,6 +25,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -51,6 +52,8 @@ TypePtr GeLUGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/gelu_grad.h b/mindspore/core/ops/grad/gelu_grad.h index bbd3bc33d6..c01f9a4d45 100644 --- a/mindspore/core/ops/grad/gelu_grad.h +++ b/mindspore/core/ops/grad/gelu_grad.h @@ -19,23 +19,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGeLUGrad = "GeLUGrad"; -class GeLUGrad : public PrimitiveC { +class GeLUGrad : public BaseOperator { public: - GeLUGrad() : PrimitiveC(kNameGeLUGrad) { InitIOName({"dy", "x", "y"}, {"z"}); } - ~GeLUGrad() = default; - MS_DECLARE_PARENT(GeLUGrad, PrimitiveC); + MIND_API_BASE_MEMBER(GeLUGrad); + GeLUGrad() : BaseOperator(kNameGeLUGrad) { InitIOName({"dy", "x", "y"}, {"z"}); } void Init() {} }; -AbstractBasePtr GeLUGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GeLUGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/grid_sampler_3d_grad.cc b/mindspore/core/ops/grad/grid_sampler_3d_grad.cc index 8909c25815..78f3f7b8ce 100644 --- a/mindspore/core/ops/grad/grid_sampler_3d_grad.cc +++ b/mindspore/core/ops/grad/grid_sampler_3d_grad.cc @@ -16,6 +16,9 @@ #include #include "ops/grad/grid_sampler_3d_grad.h" +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -90,6 +93,7 @@ TuplePtr GridSampler3DGradInferType(const PrimitivePtr &primitive, const std::ve } } // namespace +MIND_API_BASE_IMPL(GridSampler3DGrad, PrimitiveC, BaseOperator); AbstractBasePtr GridSampler3DGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/grid_sampler_3d_grad.h b/mindspore/core/ops/grad/grid_sampler_3d_grad.h index d535dd8b76..ce92f2cb74 100644 --- a/mindspore/core/ops/grad/grid_sampler_3d_grad.h +++ b/mindspore/core/ops/grad/grid_sampler_3d_grad.h @@ -21,22 +21,21 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGridSampler3DGrad = "GridSampler3DGrad"; -class GridSampler3DGrad : public PrimitiveC { +class MIND_API GridSampler3DGrad : public BaseOperator { public: - GridSampler3DGrad() : PrimitiveC(kNameGridSampler3DGrad) { InitIOName({"grad", "input_x", "grid"}, {"dx", "dgrid"}); } - ~GridSampler3DGrad() = default; - MS_DECLARE_PARENT(GridSampler3DGrad, PrimitiveC); + MIND_API_BASE_MEMBER(GridSampler3DGrad); + GridSampler3DGrad() : BaseOperator(kNameGridSampler3DGrad) { + InitIOName({"grad", "input_x", "grid"}, {"dx", "dgrid"}); + } }; -AbstractBasePtr GridSampler3DGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GridSampler3DGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimGridSampler3DGrad = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/group_conv2d_grad_input.cc b/mindspore/core/ops/grad/group_conv2d_grad_input.cc index 0a409eb682..a6840f7654 100644 --- a/mindspore/core/ops/grad/group_conv2d_grad_input.cc +++ b/mindspore/core/ops/grad/group_conv2d_grad_input.cc @@ -17,9 +17,12 @@ #include #include "ops/grad/group_conv2d_grad_input.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(GroupConv2DGradInput, PrimitiveC, BaseOperator); void GroupConv2DGradInput::Init(const int64_t &in_channel, const int64_t &out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &pad_list, const std::vector &stride, @@ -41,7 +44,7 @@ void GroupConv2DGradInput::Init(const int64_t &in_channel, const int64_t &out_ch } void GroupConv2DGradInput::set_in_channel(const int64_t &in_channel) { - (void)this->AddAttr(kInChannel, MakeValue(in_channel)); + (void)this->AddAttr(kInChannel, api::MakeValue(in_channel)); } int64_t GroupConv2DGradInput::get_in_channel() const { @@ -51,7 +54,7 @@ int64_t GroupConv2DGradInput::get_in_channel() const { } void GroupConv2DGradInput::set_out_channel(const int64_t &out_channel) { - (void)this->AddAttr(kOutChannel, MakeValue(out_channel)); + (void)this->AddAttr(kOutChannel, api::MakeValue(out_channel)); } int64_t GroupConv2DGradInput::get_out_channel() const { @@ -61,7 +64,7 @@ int64_t GroupConv2DGradInput::get_out_channel() const { } void GroupConv2DGradInput::set_kernel_size(const std::vector &kernel_size) { - (void)this->AddAttr(kKernelSize, MakeValue(kernel_size)); + (void)this->AddAttr(kKernelSize, api::MakeValue(kernel_size)); } std::vector GroupConv2DGradInput::get_kernel_size() const { @@ -72,7 +75,7 @@ std::vector GroupConv2DGradInput::get_kernel_size() const { void GroupConv2DGradInput::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode GroupConv2DGradInput::get_pad_mode() const { @@ -82,7 +85,7 @@ PadMode GroupConv2DGradInput::get_pad_mode() const { } void GroupConv2DGradInput::set_pad_list(const std::vector &pad_list) { - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } std::vector GroupConv2DGradInput::get_pad_list() const { @@ -92,7 +95,7 @@ std::vector GroupConv2DGradInput::get_pad_list() const { } void GroupConv2DGradInput::set_stride(const std::vector &stride) { - (void)this->AddAttr(kStride, MakeValue(stride)); + (void)this->AddAttr(kStride, api::MakeValue(stride)); } std::vector GroupConv2DGradInput::get_stride() const { @@ -102,7 +105,7 @@ std::vector GroupConv2DGradInput::get_stride() const { } void GroupConv2DGradInput::set_dilation(const std::vector &dilation) { - (void)this->AddAttr(kDilation, MakeValue(dilation)); + (void)this->AddAttr(kDilation, api::MakeValue(dilation)); } std::vector GroupConv2DGradInput::get_dilation() const { @@ -111,7 +114,7 @@ std::vector GroupConv2DGradInput::get_dilation() const { return GetValue>(value_ptr); } -void GroupConv2DGradInput::set_group(const int64_t &group) { (void)this->AddAttr(kGroup, MakeValue(group)); } +void GroupConv2DGradInput::set_group(const int64_t &group) { (void)this->AddAttr(kGroup, api::MakeValue(group)); } int64_t GroupConv2DGradInput::get_group() const { auto value_ptr = GetAttr(kGroup); @@ -120,7 +123,7 @@ int64_t GroupConv2DGradInput::get_group() const { } void GroupConv2DGradInput::set_input_shape(const std::vector &input_shape) { - (void)this->AddAttr(kInputShape, MakeValue(input_shape)); + (void)this->AddAttr(kInputShape, api::MakeValue(input_shape)); } std::vector GroupConv2DGradInput::get_input_shape() const { @@ -131,7 +134,7 @@ std::vector GroupConv2DGradInput::get_input_shape() const { void GroupConv2DGradInput::set_format(const Format &format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } Format GroupConv2DGradInput::get_format() const { @@ -142,7 +145,7 @@ Format GroupConv2DGradInput::get_format() const { void GroupConv2DGradInput::set_activation_type(const ActivationType &activation_type) { int64_t swi = activation_type; - (void)this->AddAttr(kActivationType, MakeValue(swi)); + (void)this->AddAttr(kActivationType, api::MakeValue(swi)); } ActivationType GroupConv2DGradInput::get_activation_type() const { @@ -151,7 +154,9 @@ ActivationType GroupConv2DGradInput::get_activation_type() const { return ActivationType(GetValue(value_ptr)); } -void GroupConv2DGradInput::set_has_bias(const bool has_bias) { (void)this->AddAttr(kHasBias, MakeValue(has_bias)); } +void GroupConv2DGradInput::set_has_bias(const bool has_bias) { + (void)this->AddAttr(kHasBias, api::MakeValue(has_bias)); +} bool GroupConv2DGradInput::get_has_bias() const { auto value_ptr = GetAttr(kHasBias); diff --git a/mindspore/core/ops/grad/group_conv2d_grad_input.h b/mindspore/core/ops/grad/group_conv2d_grad_input.h index 89fe2ca70e..60485b6b2e 100644 --- a/mindspore/core/ops/grad/group_conv2d_grad_input.h +++ b/mindspore/core/ops/grad/group_conv2d_grad_input.h @@ -19,18 +19,17 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameGroupConv2DGradInput = "GroupConv2DGradInput"; -class MS_CORE_API GroupConv2DGradInput : public PrimitiveC { +class MIND_API GroupConv2DGradInput : public BaseOperator { public: - GroupConv2DGradInput() : PrimitiveC(kNameGroupConv2DGradInput) {} - ~GroupConv2DGradInput() = default; - MS_DECLARE_PARENT(GroupConv2DGradInput, PrimitiveC); + MIND_API_BASE_MEMBER(GroupConv2DGradInput); + GroupConv2DGradInput() : BaseOperator(kNameGroupConv2DGradInput) {} void Init(const int64_t &in_channel, const int64_t &out_channel, const std::vector &kernel_size, const PadMode &pad_mode, const std::vector &pad_list, const std::vector &stride, const std::vector &dilation, const int64_t &group, const std::vector &input_shape, @@ -65,8 +64,8 @@ class MS_CORE_API GroupConv2DGradInput : public PrimitiveC { ActivationType get_activation_type() const; bool get_has_bias() const; }; -AbstractBasePtr GroupConv2DGradInputInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GroupConv2DGradInputInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/hshrink_grad.cc b/mindspore/core/ops/grad/hshrink_grad.cc index c628840b02..f5493b7ac5 100644 --- a/mindspore/core/ops/grad/hshrink_grad.cc +++ b/mindspore/core/ops/grad/hshrink_grad.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(HShrinkGrad, PrimitiveC, BaseOperator); abstract::ShapePtr HShrinkGradInferShape(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/hshrink_grad.h b/mindspore/core/ops/grad/hshrink_grad.h index 31242b8957..ae46c5b5df 100644 --- a/mindspore/core/ops/grad/hshrink_grad.h +++ b/mindspore/core/ops/grad/hshrink_grad.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_HSHRINK_GRAD_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameHShrinkGrad = "HShrinkGrad"; -class MS_CORE_API HShrinkGrad : public PrimitiveC { +class MIND_API HShrinkGrad : public BaseOperator { public: - HShrinkGrad() : PrimitiveC(kNameHShrinkGrad) { InitIOName({"gradients", "features"}, {"backprops"}); } - ~HShrinkGrad() = default; - MS_DECLARE_PARENT(HShrinkGrad, PrimitiveC); + MIND_API_BASE_MEMBER(HShrinkGrad); + HShrinkGrad() : BaseOperator(kNameHShrinkGrad) { InitIOName({"gradients", "features"}, {"backprops"}); } }; -AbstractBasePtr HShrinkGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr HShrinkGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_HSHRINK_GRAD_H_ diff --git a/mindspore/core/ops/grad/hsigmoid_grad.cc b/mindspore/core/ops/grad/hsigmoid_grad.cc index e8ecb133ab..cb80e14ea0 100644 --- a/mindspore/core/ops/grad/hsigmoid_grad.cc +++ b/mindspore/core/ops/grad/hsigmoid_grad.cc @@ -26,6 +26,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -58,6 +59,7 @@ TypePtr HSigmoidGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { return std::make_shared(HSigmoidGradInferType(primitive, input_args), diff --git a/mindspore/core/ops/grad/hsigmoid_grad.h b/mindspore/core/ops/grad/hsigmoid_grad.h index 794ea6a4df..2e960a1e75 100644 --- a/mindspore/core/ops/grad/hsigmoid_grad.h +++ b/mindspore/core/ops/grad/hsigmoid_grad.h @@ -21,22 +21,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameHSigmoidGrad = "HSigmoidGrad"; -class MS_CORE_API HSigmoidGrad : public PrimitiveC { +class MIND_API HSigmoidGrad : public BaseOperator { public: - HSigmoidGrad() : PrimitiveC(kNameHSigmoidGrad) { InitIOName({"grads", "input_x"}, {"output"}); } - ~HSigmoidGrad() = default; - MS_DECLARE_PARENT(HSigmoidGrad, PrimitiveC); + MIND_API_BASE_MEMBER(HSigmoidGrad); + HSigmoidGrad() : BaseOperator(kNameHSigmoidGrad) { InitIOName({"grads", "input_x"}, {"output"}); } }; -AbstractBasePtr HSigmoidGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr HSigmoidGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/inv_grad.cc b/mindspore/core/ops/grad/inv_grad.cc index a6046a8628..32db393693 100644 --- a/mindspore/core/ops/grad/inv_grad.cc +++ b/mindspore/core/ops/grad/inv_grad.cc @@ -19,6 +19,9 @@ #include #include "abstract/param_validator.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -55,6 +58,8 @@ TypePtr InvGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = InvGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/inv_grad.h b/mindspore/core/ops/grad/inv_grad.h index 11bfa4c281..a9842aa986 100644 --- a/mindspore/core/ops/grad/inv_grad.h +++ b/mindspore/core/ops/grad/inv_grad.h @@ -19,23 +19,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameInvGrad = "InvGrad"; -class InvGrad : public PrimitiveC { +class InvGrad : public BaseOperator { public: - InvGrad() : PrimitiveC(kNameInvGrad) { InitIOName({"x", "grad"}, {"y"}); } - ~InvGrad() = default; - MS_DECLARE_PARENT(InvGrad, PrimitiveC); + MIND_API_BASE_MEMBER(InvGrad); + InvGrad() : BaseOperator(kNameInvGrad) { InitIOName({"x", "grad"}, {"y"}); } }; -AbstractBasePtr InvGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr InvGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimInvGrad = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/layer_norm_grad.cc b/mindspore/core/ops/grad/layer_norm_grad.cc index b62a2573fa..dfa4e7c677 100644 --- a/mindspore/core/ops/grad/layer_norm_grad.cc +++ b/mindspore/core/ops/grad/layer_norm_grad.cc @@ -17,9 +17,11 @@ #include "ops/grad/layer_norm_grad.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(LayerNormGrad, PrimitiveC, BaseOperator); AbstractBasePtr LayerNormGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { // Inputs: five tensors(y_backprob, x, variance, mean, gamma). @@ -43,10 +45,10 @@ void LayerNormGrad::Init(const int64_t begin_norm_axis, const int64_t begin_para this->set_begin_params_axis(begin_params_axis); } void LayerNormGrad::set_begin_norm_axis(const int64_t begin_norm_axis) { - (void)this->AddAttr(kBeginNormAxis, MakeValue(begin_norm_axis)); + (void)this->AddAttr(kBeginNormAxis, api::MakeValue(begin_norm_axis)); } void LayerNormGrad::set_begin_params_axis(const int64_t begin_params_axis) { - (void)this->AddAttr(kBeginParamsAxis, MakeValue(begin_params_axis)); + (void)this->AddAttr(kBeginParamsAxis, api::MakeValue(begin_params_axis)); } int64_t LayerNormGrad::get_begin_norm_axis() const { auto value_ptr = this->GetAttr(kBeginNormAxis); diff --git a/mindspore/core/ops/grad/layer_norm_grad.h b/mindspore/core/ops/grad/layer_norm_grad.h index b292b9b9d5..cfe4854d95 100644 --- a/mindspore/core/ops/grad/layer_norm_grad.h +++ b/mindspore/core/ops/grad/layer_norm_grad.h @@ -20,19 +20,17 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameLayerNormGrad = prim::kLayerNormGrad; -class MS_CORE_API LayerNormGrad : public PrimitiveC { +constexpr auto kNameLayerNormGrad = "LayerNormGrad"; +class MIND_API LayerNormGrad : public BaseOperator { public: - LayerNormGrad() : PrimitiveC(kNameLayerNormGrad) {} - explicit LayerNormGrad(const std::string k_name) : PrimitiveC(k_name) {} - ~LayerNormGrad() = default; - MS_DECLARE_PARENT(LayerNormGrad, PrimitiveC); + MIND_API_BASE_MEMBER(LayerNormGrad); + LayerNormGrad() : BaseOperator(kNameLayerNormGrad) {} + explicit LayerNormGrad(const std::string k_name) : BaseOperator(k_name) {} void Init(const int64_t begin_norm_axis = 1, const int64_t begin_params_axis = 1); void set_begin_norm_axis(const int64_t begin_norm_axis); void set_begin_params_axis(const int64_t begin_params_axis); @@ -40,8 +38,8 @@ class MS_CORE_API LayerNormGrad : public PrimitiveC { int64_t get_begin_params_axis() const; }; -AbstractBasePtr LayerNormGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/log_grad.cc b/mindspore/core/ops/grad/log_grad.cc index 932a5c2e1b..b20556e9f3 100644 --- a/mindspore/core/ops/grad/log_grad.cc +++ b/mindspore/core/ops/grad/log_grad.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(LogGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLogGrad, LogGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/log_grad.h b/mindspore/core/ops/grad/log_grad.h index d4b309f6dd..1ffd44e7fe 100644 --- a/mindspore/core/ops/grad/log_grad.h +++ b/mindspore/core/ops/grad/log_grad.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogGrad = "LogGrad"; -class MS_CORE_API LogGrad : public PrimitiveC { +class MIND_API LogGrad : public BaseOperator { public: - LogGrad() : PrimitiveC(kNameLogGrad) {} - ~LogGrad() = default; - MS_DECLARE_PARENT(LogGrad, PrimitiveC); + MIND_API_BASE_MEMBER(LogGrad); + LogGrad() : BaseOperator(kNameLogGrad) {} void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/log_softmax_grad.cc b/mindspore/core/ops/grad/log_softmax_grad.cc index b75a529c6b..951e005a1a 100644 --- a/mindspore/core/ops/grad/log_softmax_grad.cc +++ b/mindspore/core/ops/grad/log_softmax_grad.cc @@ -17,6 +17,7 @@ #include "ops/grad/log_softmax_grad.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -51,9 +52,10 @@ TypePtr LogSoftmaxGradInferType(const PrimitivePtr &prim, const std::vectorset_axis(axis); } -void LogSoftmaxGrad::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void LogSoftmaxGrad::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t LogSoftmaxGrad::get_axis() const { auto value_ptr = GetAttr(kAxis); diff --git a/mindspore/core/ops/grad/log_softmax_grad.h b/mindspore/core/ops/grad/log_softmax_grad.h index 4b8be37de2..bb48332a51 100644 --- a/mindspore/core/ops/grad/log_softmax_grad.h +++ b/mindspore/core/ops/grad/log_softmax_grad.h @@ -20,26 +20,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogSoftmaxGrad = "LogSoftmaxGrad"; -class LogSoftmaxGrad : public PrimitiveC { +class MIND_API LogSoftmaxGrad : public BaseOperator { public: - LogSoftmaxGrad() : PrimitiveC(kNameLogSoftmaxGrad) { InitIOName({"x", "grad"}, {"y"}); } - explicit LogSoftmaxGrad(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "grad"}, {"y"}); } - ~LogSoftmaxGrad() = default; - MS_DECLARE_PARENT(LogSoftmaxGrad, PrimitiveC); + MIND_API_BASE_MEMBER(LogSoftmaxGrad); + LogSoftmaxGrad() : BaseOperator(kNameLogSoftmaxGrad) { InitIOName({"x", "grad"}, {"y"}); } + explicit LogSoftmaxGrad(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "grad"}, {"y"}); } void Init(const int64_t axis = -1); void set_axis(const int64_t epsilon); int64_t get_axis() const; }; -AbstractBasePtr LogSoftmaxGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogSoftmaxGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/lstm_grad.cc b/mindspore/core/ops/grad/lstm_grad.cc index 3dd6ec2d06..b8305ebecc 100644 --- a/mindspore/core/ops/grad/lstm_grad.cc +++ b/mindspore/core/ops/grad/lstm_grad.cc @@ -17,51 +17,57 @@ #include "ops/grad/lstm_grad.h" #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void LSTMGrad::set_input_size(const int64_t input_size) { (void)CheckAndConvertUtils::CheckInteger(kInput_size, input_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kInput_size, MakeValue(input_size)); + (void)AddAttr(kInput_size, api::MakeValue(input_size)); } int64_t LSTMGrad::get_input_size() const { return GetValue(GetAttr(kInput_size)); } void LSTMGrad::set_hidden_size(const int64_t hidden_size) { (void)CheckAndConvertUtils::CheckInteger(kHidden_size, hidden_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kHidden_size, MakeValue(hidden_size)); + (void)AddAttr(kHidden_size, api::MakeValue(hidden_size)); } int64_t LSTMGrad::get_hidden_size() const { return GetValue(GetAttr(kHidden_size)); } void LSTMGrad::set_num_layers(const int64_t num_layers) { (void)CheckAndConvertUtils::CheckInteger(kNumLayers, num_layers, kGreaterThan, 0, this->name()); - (void)AddAttr(kNumLayers, MakeValue(num_layers)); + (void)AddAttr(kNumLayers, api::MakeValue(num_layers)); } int64_t LSTMGrad::get_num_layers() const { return GetValue(GetAttr(kNumLayers)); } -void LSTMGrad::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, MakeValue(has_bias)); } +void LSTMGrad::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, api::MakeValue(has_bias)); } bool LSTMGrad::get_has_bias() const { auto value_ptr = this->GetAttr(kHasBias); return GetValue(value_ptr); } void LSTMGrad::set_dropout(const float dropout) { (void)CheckAndConvertUtils::CheckInRange(kDropout, dropout, kIncludeBoth, {0.0, 1.0}, this->name()); - (void)AddAttr(kDropout, MakeValue(dropout)); + (void)AddAttr(kDropout, api::MakeValue(dropout)); } float LSTMGrad::get_dropout() const { auto value_ptr = this->GetAttr(kDropout); return GetValue(value_ptr); } -void LSTMGrad::set_bidirectional(const bool bidirectional) { (void)AddAttr(kBidirectional, MakeValue(bidirectional)); } +void LSTMGrad::set_bidirectional(const bool bidirectional) { + (void)AddAttr(kBidirectional, api::MakeValue(bidirectional)); +} bool LSTMGrad::get_bidirectional() const { auto value_ptr = this->GetAttr(kBidirectional); return GetValue(value_ptr); } void LSTMGrad::set_num_directions(const int64_t num_directions) { - (void)AddAttr(kNumDirections, MakeValue(num_directions)); + (void)AddAttr(kNumDirections, api::MakeValue(num_directions)); } int64_t LSTMGrad::get_num_directions() const { return GetValue(GetAttr(kNumDirections)); } -void LSTMGrad::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, MakeValue(zoneout_cell)); } +void LSTMGrad::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, api::MakeValue(zoneout_cell)); } float LSTMGrad::get_zoneout_cell() const { return GetValue(this->GetAttr(kZoneoutCell)); } -void LSTMGrad::set_zoneout_hidden(float zoneout_hidden) { (void)AddAttr(kZoneoutHidden, MakeValue(zoneout_hidden)); } +void LSTMGrad::set_zoneout_hidden(float zoneout_hidden) { + (void)AddAttr(kZoneoutHidden, api::MakeValue(zoneout_hidden)); +} float LSTMGrad::get_zoneout_hidden() const { return GetValue(this->GetAttr(kZoneoutHidden)); } @@ -84,6 +90,7 @@ void LSTMGrad::Init(const int64_t input_size, const int64_t hidden_size, const i this->set_zoneout_hidden(zoneout_hidden); } +MIND_API_BASE_IMPL(LSTMGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLSTMGrad, LSTMGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/lstm_grad.h b/mindspore/core/ops/grad/lstm_grad.h index 4c8787b8df..01530f1e39 100644 --- a/mindspore/core/ops/grad/lstm_grad.h +++ b/mindspore/core/ops/grad/lstm_grad.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLSTMGrad = "LSTMGrad"; -class MS_CORE_API LSTMGrad : public PrimitiveC { +class MIND_API LSTMGrad : public BaseOperator { public: - LSTMGrad() : PrimitiveC(kNameLSTMGrad) {} - ~LSTMGrad() = default; - MS_DECLARE_PARENT(LSTMGrad, PrimitiveC); + MIND_API_BASE_MEMBER(LSTMGrad); + LSTMGrad() : BaseOperator(kNameLSTMGrad) {} void Init(const int64_t input_size, const int64_t hidden_size, const int64_t num_layers, const bool has_bias, const float dropout, const bool bidirectional = false, const float zoneout_cell = 0.0f, const float zoneout_hidden = 0.0f); @@ -55,8 +53,8 @@ class MS_CORE_API LSTMGrad : public PrimitiveC { float get_zoneout_hidden() const; int64_t get_good_ld(const int64_t dim, const int64_t type_size); }; -AbstractBasePtr LstmGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LstmGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/lstm_grad_data.cc b/mindspore/core/ops/grad/lstm_grad_data.cc index f5f26a8513..ef56021845 100644 --- a/mindspore/core/ops/grad/lstm_grad_data.cc +++ b/mindspore/core/ops/grad/lstm_grad_data.cc @@ -17,54 +17,56 @@ #include "ops/grad/lstm_grad_data.h" #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void LSTMGradData::set_input_size(const int64_t input_size) { (void)CheckAndConvertUtils::CheckInteger(kInput_size, input_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kInput_size, MakeValue(input_size)); + (void)AddAttr(kInput_size, api::MakeValue(input_size)); } int64_t LSTMGradData::get_input_size() const { return GetValue(GetAttr(kInput_size)); } void LSTMGradData::set_hidden_size(const int64_t hidden_size) { (void)CheckAndConvertUtils::CheckInteger(kHidden_size, hidden_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kHidden_size, MakeValue(hidden_size)); + (void)AddAttr(kHidden_size, api::MakeValue(hidden_size)); } int64_t LSTMGradData::get_hidden_size() const { return GetValue(GetAttr(kHidden_size)); } void LSTMGradData::set_num_layers(const int64_t num_layers) { (void)CheckAndConvertUtils::CheckInteger(kNumLayers, num_layers, kGreaterThan, 0, this->name()); - (void)AddAttr(kNumLayers, MakeValue(num_layers)); + (void)AddAttr(kNumLayers, api::MakeValue(num_layers)); } int64_t LSTMGradData::get_num_layers() const { return GetValue(GetAttr(kNumLayers)); } -void LSTMGradData::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, MakeValue(has_bias)); } +void LSTMGradData::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, api::MakeValue(has_bias)); } bool LSTMGradData::get_has_bias() const { auto value_ptr = this->GetAttr(kHasBias); return GetValue(value_ptr); } void LSTMGradData::set_dropout(const float dropout) { (void)CheckAndConvertUtils::CheckInRange(kDropout, dropout, kIncludeBoth, {0.0, 1.0}, this->name()); - (void)AddAttr(kDropout, MakeValue(dropout)); + (void)AddAttr(kDropout, api::MakeValue(dropout)); } float LSTMGradData::get_dropout() const { auto value_ptr = this->GetAttr(kDropout); return GetValue(value_ptr); } void LSTMGradData::set_bidirectional(const bool bidirectional) { - (void)AddAttr(kBidirectional, MakeValue(bidirectional)); + (void)AddAttr(kBidirectional, api::MakeValue(bidirectional)); } bool LSTMGradData::get_bidirectional() const { auto value_ptr = this->GetAttr(kBidirectional); return GetValue(value_ptr); } void LSTMGradData::set_num_directions(const int64_t num_directions) { - (void)AddAttr(kNumDirections, MakeValue(num_directions)); + (void)AddAttr(kNumDirections, api::MakeValue(num_directions)); } int64_t LSTMGradData::get_num_directions() const { return GetValue(GetAttr(kNumDirections)); } -void LSTMGradData::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, MakeValue(zoneout_cell)); } +void LSTMGradData::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, api::MakeValue(zoneout_cell)); } float LSTMGradData::get_zoneout_cell() const { return GetValue(this->GetAttr(kZoneoutCell)); } void LSTMGradData::set_zoneout_hidden(float zoneout_hidden) { - (void)AddAttr(kZoneoutHidden, MakeValue(zoneout_hidden)); + (void)AddAttr(kZoneoutHidden, api::MakeValue(zoneout_hidden)); } float LSTMGradData::get_zoneout_hidden() const { return GetValue(this->GetAttr(kZoneoutHidden)); } @@ -88,6 +90,7 @@ void LSTMGradData::Init(const int64_t input_size, const int64_t hidden_size, con this->set_zoneout_hidden(zoneout_hidden); } +MIND_API_BASE_IMPL(LSTMGradData, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLSTMGradData, LSTMGradData); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/lstm_grad_data.h b/mindspore/core/ops/grad/lstm_grad_data.h index 9d361f98cf..2273af8135 100644 --- a/mindspore/core/ops/grad/lstm_grad_data.h +++ b/mindspore/core/ops/grad/lstm_grad_data.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLSTMGradData = "LSTMGradData"; -class MS_CORE_API LSTMGradData : public PrimitiveC { +class MIND_API LSTMGradData : public BaseOperator { public: - LSTMGradData() : PrimitiveC(kNameLSTMGradData) {} - ~LSTMGradData() = default; - MS_DECLARE_PARENT(LSTMGradData, PrimitiveC); + MIND_API_BASE_MEMBER(LSTMGradData); + LSTMGradData() : BaseOperator(kNameLSTMGradData) {} void Init(const int64_t input_size, const int64_t hidden_size, const int64_t num_layers, const bool has_bias, const float dropout, const bool bidirectional = false, const float zoneout_cell = 0.0f, const float zoneout_hidden = 0.0f); @@ -55,8 +53,8 @@ class MS_CORE_API LSTMGradData : public PrimitiveC { float get_zoneout_hidden() const; int64_t get_good_ld(const int64_t dim, const int64_t type_size); }; -AbstractBasePtr LstmGradDataInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LstmGradDataInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/lstm_grad_weight.cc b/mindspore/core/ops/grad/lstm_grad_weight.cc index 6794ce0ba1..05f7c5d859 100644 --- a/mindspore/core/ops/grad/lstm_grad_weight.cc +++ b/mindspore/core/ops/grad/lstm_grad_weight.cc @@ -17,54 +17,56 @@ #include "ops/grad/lstm_grad_weight.h" #include #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void LSTMGradWeight::set_input_size(const int64_t input_size) { (void)CheckAndConvertUtils::CheckInteger(kInput_size, input_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kInput_size, MakeValue(input_size)); + (void)AddAttr(kInput_size, api::MakeValue(input_size)); } int64_t LSTMGradWeight::get_input_size() const { return GetValue(GetAttr(kInput_size)); } void LSTMGradWeight::set_hidden_size(const int64_t hidden_size) { (void)CheckAndConvertUtils::CheckInteger(kHidden_size, hidden_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kHidden_size, MakeValue(hidden_size)); + (void)AddAttr(kHidden_size, api::MakeValue(hidden_size)); } int64_t LSTMGradWeight::get_hidden_size() const { return GetValue(GetAttr(kHidden_size)); } void LSTMGradWeight::set_num_layers(const int64_t num_layers) { (void)CheckAndConvertUtils::CheckInteger(kNumLayers, num_layers, kGreaterThan, 0, this->name()); - (void)AddAttr(kNumLayers, MakeValue(num_layers)); + (void)AddAttr(kNumLayers, api::MakeValue(num_layers)); } int64_t LSTMGradWeight::get_num_layers() const { return GetValue(GetAttr(kNumLayers)); } -void LSTMGradWeight::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, MakeValue(has_bias)); } +void LSTMGradWeight::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, api::MakeValue(has_bias)); } bool LSTMGradWeight::get_has_bias() const { auto value_ptr = this->GetAttr(kHasBias); return GetValue(value_ptr); } void LSTMGradWeight::set_dropout(const float dropout) { (void)CheckAndConvertUtils::CheckInRange(kDropout, dropout, kIncludeBoth, {0.0, 1.0}, this->name()); - (void)AddAttr(kDropout, MakeValue(dropout)); + (void)AddAttr(kDropout, api::MakeValue(dropout)); } float LSTMGradWeight::get_dropout() const { auto value_ptr = this->GetAttr(kDropout); return GetValue(value_ptr); } void LSTMGradWeight::set_bidirectional(const bool bidirectional) { - (void)AddAttr(kBidirectional, MakeValue(bidirectional)); + (void)AddAttr(kBidirectional, api::MakeValue(bidirectional)); } bool LSTMGradWeight::get_bidirectional() const { auto value_ptr = this->GetAttr(kBidirectional); return GetValue(value_ptr); } void LSTMGradWeight::set_num_directions(const int64_t num_directions) { - (void)AddAttr(kNumDirections, MakeValue(num_directions)); + (void)AddAttr(kNumDirections, api::MakeValue(num_directions)); } int64_t LSTMGradWeight::get_num_directions() const { return GetValue(GetAttr(kNumDirections)); } -void LSTMGradWeight::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, MakeValue(zoneout_cell)); } +void LSTMGradWeight::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, api::MakeValue(zoneout_cell)); } float LSTMGradWeight::get_zoneout_cell() const { return GetValue(this->GetAttr(kZoneoutCell)); } void LSTMGradWeight::set_zoneout_hidden(float zoneout_hidden) { - (void)AddAttr(kZoneoutHidden, MakeValue(zoneout_hidden)); + (void)AddAttr(kZoneoutHidden, api::MakeValue(zoneout_hidden)); } float LSTMGradWeight::get_zoneout_hidden() const { return GetValue(this->GetAttr(kZoneoutHidden)); } @@ -88,6 +90,7 @@ void LSTMGradWeight::Init(const int64_t input_size, const int64_t hidden_size, c this->set_zoneout_hidden(zoneout_hidden); } +MIND_API_BASE_IMPL(LSTMGradWeight, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLSTMGradWeight, LSTMGradWeight); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/lstm_grad_weight.h b/mindspore/core/ops/grad/lstm_grad_weight.h index ec2cb53b7b..7c7b79965c 100644 --- a/mindspore/core/ops/grad/lstm_grad_weight.h +++ b/mindspore/core/ops/grad/lstm_grad_weight.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLSTMGradWeight = "LSTMGradWeight"; -class MS_CORE_API LSTMGradWeight : public PrimitiveC { +class MIND_API LSTMGradWeight : public BaseOperator { public: - LSTMGradWeight() : PrimitiveC(kNameLSTMGradWeight) {} - ~LSTMGradWeight() = default; - MS_DECLARE_PARENT(LSTMGradWeight, PrimitiveC); + MIND_API_BASE_MEMBER(LSTMGradWeight); + LSTMGradWeight() : BaseOperator(kNameLSTMGradWeight) {} void Init(const int64_t input_size, const int64_t hidden_size, const int64_t num_layers, const bool has_bias, const float dropout, const bool bidirectional = false, const float zoneout_cell = 0.0f, const float zoneout_hidden = 0.0f); @@ -55,8 +53,8 @@ class MS_CORE_API LSTMGradWeight : public PrimitiveC { float get_zoneout_hidden() const; int64_t get_good_ld(const int64_t dim, const int64_t type_size); }; -AbstractBasePtr LstmGradWeightInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LstmGradWeightInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/max_pool_grad.cc b/mindspore/core/ops/grad/max_pool_grad.cc index 68da0b008b..ef742d4caf 100644 --- a/mindspore/core/ops/grad/max_pool_grad.cc +++ b/mindspore/core/ops/grad/max_pool_grad.cc @@ -16,9 +16,12 @@ #include "ops/grad/max_pool_grad.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(MaxPoolGrad, PrimitiveC, PoolGrad); REGISTER_PRIMITIVE_C(kNameMaxPoolGrad, MaxPoolGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/max_pool_grad.h b/mindspore/core/ops/grad/max_pool_grad.h index 3cefec3db0..c63537c127 100644 --- a/mindspore/core/ops/grad/max_pool_grad.h +++ b/mindspore/core/ops/grad/max_pool_grad.h @@ -21,22 +21,20 @@ #include #include #include "ops/grad/pool_grad.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMaxPoolGrad = "MaxPoolGrad"; -class MS_CORE_API MaxPoolGrad : public PoolGrad { +class MIND_API MaxPoolGrad : public PoolGrad { public: + MIND_API_BASE_MEMBER(MaxPoolGrad); MaxPoolGrad() : PoolGrad(kNameMaxPoolGrad) { InitIOName({"x_origin", "out_origin", "grad"}, {"output"}); } - ~MaxPoolGrad() = default; - MS_DECLARE_PARENT(MaxPoolGrad, PoolGrad); }; -AbstractBasePtr MaxPoolGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MaxPoolGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimMaxPoolGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/maximum_grad.cc b/mindspore/core/ops/grad/maximum_grad.cc index 3115514218..90182c8ec3 100644 --- a/mindspore/core/ops/grad/maximum_grad.cc +++ b/mindspore/core/ops/grad/maximum_grad.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,8 +61,10 @@ TuplePtr MaximumGradInferType(const PrimitivePtr &primitive, const std::vector(type_tuple); } } // namespace -AbstractBasePtr MaximumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args) { + +MIND_API_BASE_IMPL(MaximumGrad, PrimitiveC, BaseOperator); +abstract::AbstractBasePtr MaximumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); auto infer_type = MaximumGradInferType(primitive, input_args); auto infer_shape = MaximumGradInferShape(primitive, input_args); @@ -72,9 +75,9 @@ void MaximumGrad::Init(const bool grad_x, const bool grad_y) { set_grad_y(grad_y); } -void MaximumGrad::set_grad_x(const bool grad_x) { (void)this->AddAttr(kGradX, MakeValue(grad_x)); } +void MaximumGrad::set_grad_x(const bool grad_x) { (void)this->AddAttr(kGradX, api::MakeValue(grad_x)); } -void MaximumGrad::set_grad_y(const bool grad_y) { (void)this->AddAttr(kGradY, MakeValue(grad_y)); } +void MaximumGrad::set_grad_y(const bool grad_y) { (void)this->AddAttr(kGradY, api::MakeValue(grad_y)); } bool MaximumGrad::get_grad_x() const { auto value_ptr = GetAttr(kGradX); diff --git a/mindspore/core/ops/grad/maximum_grad.h b/mindspore/core/ops/grad/maximum_grad.h index 05f030caf1..1a525f6ec1 100644 --- a/mindspore/core/ops/grad/maximum_grad.h +++ b/mindspore/core/ops/grad/maximum_grad.h @@ -18,26 +18,24 @@ #define MINDSPORE_CORE_OPS_MAXIMUM_GRAD_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMaximumGrad = "MaximumGrad"; -class MS_CORE_API MaximumGrad : public PrimitiveC { +class MIND_API MaximumGrad : public BaseOperator { public: - MaximumGrad() : PrimitiveC(kNameMaximumGrad) { InitIOName({"x1", "x2", "grads"}, {"y1", "y2"}); } - ~MaximumGrad() = default; - MS_DECLARE_PARENT(MaximumGrad, PrimitiveC); + MIND_API_BASE_MEMBER(MaximumGrad); + MaximumGrad() : BaseOperator(kNameMaximumGrad) { InitIOName({"x1", "x2", "grads"}, {"y1", "y2"}); } void Init(const bool grad_x = true, const bool grad_y = true); void set_grad_x(const bool grad_x); void set_grad_y(const bool grad_y); bool get_grad_x() const; bool get_grad_y() const; }; -AbstractBasePtr MaximumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MaximumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimMaximumGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/minimum_grad.cc b/mindspore/core/ops/grad/minimum_grad.cc index cf97554004..5ecc132de8 100644 --- a/mindspore/core/ops/grad/minimum_grad.cc +++ b/mindspore/core/ops/grad/minimum_grad.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,8 +61,10 @@ TuplePtr MinimumGradInferType(const PrimitivePtr &primitive, const std::vector(type_tuple); } } // namespace -AbstractBasePtr MinimumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args) { + +MIND_API_BASE_IMPL(MinimumGrad, PrimitiveC, BaseOperator); +abstract::AbstractBasePtr MinimumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); auto infer_type = MinimumGradInferType(primitive, input_args); auto infer_shape = MinimumGradInferShape(primitive, input_args); @@ -72,9 +75,9 @@ void MinimumGrad::Init(const bool grad_x, const bool grad_y) { set_grad_y(grad_y); } -void MinimumGrad::set_grad_x(const bool grad_x) { (void)this->AddAttr(kGradX, MakeValue(grad_x)); } +void MinimumGrad::set_grad_x(const bool grad_x) { (void)this->AddAttr(kGradX, api::MakeValue(grad_x)); } -void MinimumGrad::set_grad_y(const bool grad_y) { (void)this->AddAttr(kGradY, MakeValue(grad_y)); } +void MinimumGrad::set_grad_y(const bool grad_y) { (void)this->AddAttr(kGradY, api::MakeValue(grad_y)); } bool MinimumGrad::get_grad_x() const { auto value_ptr = GetAttr(kGradX); diff --git a/mindspore/core/ops/grad/minimum_grad.h b/mindspore/core/ops/grad/minimum_grad.h index 1c495d08ca..65eef54bb8 100644 --- a/mindspore/core/ops/grad/minimum_grad.h +++ b/mindspore/core/ops/grad/minimum_grad.h @@ -18,26 +18,24 @@ #define MINDSPORE_CORE_OPS_MINIMUM_GRAD_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMinimumGrad = "MinimumGrad"; -class MS_CORE_API MinimumGrad : public PrimitiveC { +class MIND_API MinimumGrad : public BaseOperator { public: - MinimumGrad() : PrimitiveC(kNameMinimumGrad) { InitIOName({"x1", "x2", "grads"}, {"y1", "y2"}); } - ~MinimumGrad() = default; - MS_DECLARE_PARENT(MinimumGrad, PrimitiveC); + MIND_API_BASE_MEMBER(MinimumGrad); + MinimumGrad() : BaseOperator(kNameMinimumGrad) { InitIOName({"x1", "x2", "grads"}, {"y1", "y2"}); } void Init(const bool grad_x = true, const bool grad_y = true); void set_grad_x(const bool grad_x); void set_grad_y(const bool grad_y); bool get_grad_x() const; bool get_grad_y() const; }; -AbstractBasePtr MinimumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MinimumGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimMinimumGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/mul_grad.cc b/mindspore/core/ops/grad/mul_grad.cc index ad4be6d894..f1dd7ce10d 100644 --- a/mindspore/core/ops/grad/mul_grad.cc +++ b/mindspore/core/ops/grad/mul_grad.cc @@ -17,9 +17,11 @@ #include "ops/grad/mul_grad.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(MulGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMulGrad, MulGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/mul_grad.h b/mindspore/core/ops/grad/mul_grad.h index 993ec9ad93..8b2571738c 100644 --- a/mindspore/core/ops/grad/mul_grad.h +++ b/mindspore/core/ops/grad/mul_grad.h @@ -17,18 +17,16 @@ #define MINDSPORE_CORE_OPS_MUL_GRAD_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMulGrad = "MulGrad"; -class MS_CORE_API MulGrad : public PrimitiveC { +class MIND_API MulGrad : public BaseOperator { public: - MulGrad() : PrimitiveC(kNameMulGrad) {} - ~MulGrad() = default; - MS_DECLARE_PARENT(MulGrad, PrimitiveC); + MIND_API_BASE_MEMBER(MulGrad); + MulGrad() : BaseOperator(kNameMulGrad) {} void Init() const {} }; diff --git a/mindspore/core/ops/grad/neg_grad.cc b/mindspore/core/ops/grad/neg_grad.cc index b5fe0feeb9..13617432ae 100644 --- a/mindspore/core/ops/grad/neg_grad.cc +++ b/mindspore/core/ops/grad/neg_grad.cc @@ -15,9 +15,12 @@ */ #include "ops/grad/neg_grad.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(NegGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameNegGrad, NegGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/neg_grad.h b/mindspore/core/ops/grad/neg_grad.h index 5f93dc6967..107ea4b3c2 100644 --- a/mindspore/core/ops/grad/neg_grad.h +++ b/mindspore/core/ops/grad/neg_grad.h @@ -22,20 +22,16 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNegGrad = "NegGrad"; -class MS_CORE_API NegGrad : public PrimitiveC { +class MIND_API NegGrad : public BaseOperator { public: - NegGrad() : PrimitiveC(kNameNegGrad) {} - ~NegGrad() = default; - MS_DECLARE_PARENT(NegGrad, PrimitiveC); + MIND_API_BASE_MEMBER(NegGrad); + NegGrad() : BaseOperator(kNameNegGrad) {} void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/nllloss_grad.cc b/mindspore/core/ops/grad/nllloss_grad.cc index e3cdb336fd..2e069cca5b 100644 --- a/mindspore/core/ops/grad/nllloss_grad.cc +++ b/mindspore/core/ops/grad/nllloss_grad.cc @@ -17,6 +17,7 @@ #include "ops/grad/nllloss_grad.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -24,7 +25,7 @@ void NLLLossGrad::Init(const Reduction &reduction) { set_reduction(reduction); } void NLLLossGrad::set_reduction(const Reduction &reduction) { int64_t reduce = reduction; - (void)AddAttr(kReduction, MakeValue(reduce)); + (void)AddAttr(kReduction, api::MakeValue(reduce)); } Reduction NLLLossGrad::get_reduction() const { @@ -32,6 +33,7 @@ Reduction NLLLossGrad::get_reduction() const { return Reduction(GetValue(value_ptr)); } +MIND_API_BASE_IMPL(NLLLossGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameNLLLossGrad, NLLLossGrad) } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/nllloss_grad.h b/mindspore/core/ops/grad/nllloss_grad.h index 5befba8ac3..d675f7abc4 100644 --- a/mindspore/core/ops/grad/nllloss_grad.h +++ b/mindspore/core/ops/grad/nllloss_grad.h @@ -19,25 +19,21 @@ #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" #include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNLLLossGrad = "NLLLossGrad"; /// \brief NLLLossGrad operation. Refer to Python API @ref mindspore.ops.NLLLossGrad for more details. -class MS_CORE_API NLLLossGrad : public PrimitiveC { +class MIND_API NLLLossGrad : public BaseOperator { public: + MIND_API_BASE_MEMBER(NLLLossGrad); /// \brief Constructor. - NLLLossGrad() : PrimitiveC(kNameNLLLossGrad) { + NLLLossGrad() : BaseOperator(kNameNLLLossGrad) { InitIOName({"logits", "loss_grad", "labels", "weight", "total_weight"}, {"logits_grad"}); } - /// \brief Destructor. - ~NLLLossGrad() = default; - - MS_DECLARE_PARENT(NLLLossGrad, PrimitiveC); - /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.NLLLossGrad for the inputs. void Init(const Reduction &reduction = NONE); diff --git a/mindspore/core/ops/grad/pool_grad.cc b/mindspore/core/ops/grad/pool_grad.cc index 8b07c09e40..6c0769d818 100644 --- a/mindspore/core/ops/grad/pool_grad.cc +++ b/mindspore/core/ops/grad/pool_grad.cc @@ -16,9 +16,11 @@ #include "ops/grad/pool_grad.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PoolGrad, PrimitiveC, BaseOperator); std::vector PoolGrad::_grad_check_vector(const std::string &arg_name, std::vector arg_val, const std::string &op_name) { std::vector ret; @@ -56,22 +58,22 @@ void PoolGrad::Init(const std::vector &kernel_size, const std::vector &kernel_size) { std::vector k_size = _grad_check_vector(kKernelSize, kernel_size, this->name()); - (void)this->AddAttr(kKernelSize, MakeValue(k_size)); + (void)this->AddAttr(kKernelSize, api::MakeValue(k_size)); } void PoolGrad::set_strides(const std::vector &strides) { std::vector strides_ = _grad_check_vector(kStrides, strides, this->name()); - (void)this->AddAttr(kStrides, MakeValue(strides_)); + (void)this->AddAttr(kStrides, api::MakeValue(strides_)); } void PoolGrad::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } void PoolGrad::set_format(const Format &format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } std::vector PoolGrad::get_kernel_size() const { diff --git a/mindspore/core/ops/grad/pool_grad.h b/mindspore/core/ops/grad/pool_grad.h index 3ceb81927d..1173c6b88c 100644 --- a/mindspore/core/ops/grad/pool_grad.h +++ b/mindspore/core/ops/grad/pool_grad.h @@ -20,21 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNamePoolGrad = "PoolGrad"; -class MS_CORE_API PoolGrad : public PrimitiveC { +class MIND_API PoolGrad : public BaseOperator { public: - PoolGrad() : PrimitiveC(kNamePoolGrad) { InitIOName({"x_origin", "out_origin", "grad"}, {"output"}); } - explicit PoolGrad(const std::string k_name) : PrimitiveC(k_name) { + MIND_API_BASE_MEMBER(PoolGrad); + PoolGrad() : BaseOperator(kNamePoolGrad) { InitIOName({"x_origin", "out_origin", "grad"}, {"output"}); } + explicit PoolGrad(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x_origin", "out_origin", "grad"}, {"output"}); } - ~PoolGrad() = default; - MS_DECLARE_PARENT(PoolGrad, PrimitiveC); virtual void Init(const std::vector &kernel_size = {1}, const std::vector &strides = {1}, const PadMode &pad_mode = VALID, const Format &format = NCHW); virtual void set_kernel_size(const std::vector &kernel_size); diff --git a/mindspore/core/ops/grad/pooling_grad.cc b/mindspore/core/ops/grad/pooling_grad.cc index 4375197077..99276459f9 100644 --- a/mindspore/core/ops/grad/pooling_grad.cc +++ b/mindspore/core/ops/grad/pooling_grad.cc @@ -16,9 +16,11 @@ #include "ops/grad/pooling_grad.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PoolingGrad, PrimitiveC, BaseOperator); void PoolingGrad::Init(const PoolMode &pool_mode, const std::vector &window, const std::vector &stride, const PadMode &pad_mode, const std::vector &pad_list, const RoundMode &round_mode, const Format &format, @@ -35,7 +37,7 @@ void PoolingGrad::Init(const PoolMode &pool_mode, const std::vector &wi void PoolingGrad::set_pool_mode(const PoolMode &pool_mode) { int64_t swi = pool_mode; - (void)this->AddAttr(kPoolMode, MakeValue(swi)); + (void)this->AddAttr(kPoolMode, api::MakeValue(swi)); } PoolMode PoolingGrad::get_pool_mode() const { @@ -43,7 +45,9 @@ PoolMode PoolingGrad::get_pool_mode() const { return PoolMode(GetValue(value_ptr)); } -void PoolingGrad::set_window(const std::vector &window) { (void)this->AddAttr(kWindow, MakeValue(window)); } +void PoolingGrad::set_window(const std::vector &window) { + (void)this->AddAttr(kWindow, api::MakeValue(window)); +} std::vector PoolingGrad::get_window() const { auto value_ptr = GetAttr(kWindow); @@ -51,7 +55,9 @@ std::vector PoolingGrad::get_window() const { return GetValue>(value_ptr); } -void PoolingGrad::set_stride(const std::vector &stride) { (void)this->AddAttr(kStride, MakeValue(stride)); } +void PoolingGrad::set_stride(const std::vector &stride) { + (void)this->AddAttr(kStride, api::MakeValue(stride)); +} std::vector PoolingGrad::get_stride() const { auto value_ptr = GetAttr(kStride); @@ -61,7 +67,7 @@ std::vector PoolingGrad::get_stride() const { void PoolingGrad::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode PoolingGrad::get_pad_mode() const { @@ -71,7 +77,7 @@ PadMode PoolingGrad::get_pad_mode() const { } void PoolingGrad::set_pad_list(const std::vector &pad_list) { - (void)this->AddAttr(kPadList, MakeValue(pad_list)); + (void)this->AddAttr(kPadList, api::MakeValue(pad_list)); } std::vector PoolingGrad::get_pad_list() const { @@ -82,7 +88,7 @@ std::vector PoolingGrad::get_pad_list() const { void PoolingGrad::set_round_mode(const RoundMode &round_mode) { int64_t swi = round_mode; - (void)this->AddAttr(kRoundMode, MakeValue(swi)); + (void)this->AddAttr(kRoundMode, api::MakeValue(swi)); } RoundMode PoolingGrad::get_round_mode() const { @@ -93,7 +99,7 @@ RoundMode PoolingGrad::get_round_mode() const { void PoolingGrad::set_format(const Format &format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } Format PoolingGrad::get_format() const { @@ -102,7 +108,7 @@ Format PoolingGrad::get_format() const { return Format(GetValue(value_ptr)); } -void PoolingGrad::set_global(const bool global) { (void)this->AddAttr(kGlobal, MakeValue(global)); } +void PoolingGrad::set_global(const bool global) { (void)this->AddAttr(kGlobal, api::MakeValue(global)); } bool PoolingGrad::get_global() const { auto value_ptr = GetAttr(kGlobal); diff --git a/mindspore/core/ops/grad/pooling_grad.h b/mindspore/core/ops/grad/pooling_grad.h index b54feee1e2..8e2bf9f4f4 100644 --- a/mindspore/core/ops/grad/pooling_grad.h +++ b/mindspore/core/ops/grad/pooling_grad.h @@ -20,18 +20,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNamePoolingGrad = "PoolingGrad"; -class MS_CORE_API PoolingGrad : public PrimitiveC { +class MIND_API PoolingGrad : public BaseOperator { public: - PoolingGrad() : PrimitiveC(kNamePoolingGrad) {} - ~PoolingGrad() = default; - MS_DECLARE_PARENT(PoolingGrad, PrimitiveC); + MIND_API_BASE_MEMBER(PoolingGrad); + PoolingGrad() : BaseOperator(kNamePoolingGrad) {} void Init(const PoolMode &pool_mode, const std::vector &window, const std::vector &stride, const PadMode &pad_mode, const std::vector &pad_list, const RoundMode &round_mode, const Format &format = NCHW, const bool global = false); diff --git a/mindspore/core/ops/grad/power_grad.cc b/mindspore/core/ops/grad/power_grad.cc index 24e6854949..a0c5d89a15 100644 --- a/mindspore/core/ops/grad/power_grad.cc +++ b/mindspore/core/ops/grad/power_grad.cc @@ -23,24 +23,26 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void PowerGrad::set_power(const float power) { (void)this->AddAttr(kPower, MakeValue(power)); } +MIND_API_BASE_IMPL(PowerGrad, PrimitiveC, BaseOperator); +void PowerGrad::set_power(const float power) { (void)this->AddAttr(kPower, api::MakeValue(power)); } float PowerGrad::get_power() const { auto value_ptr = GetAttr(kPower); MS_EXCEPTION_IF_NULL(value_ptr); return GetValue(value_ptr); } -void PowerGrad::set_scale(const float scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } +void PowerGrad::set_scale(const float scale) { (void)this->AddAttr(kScale, api::MakeValue(scale)); } float PowerGrad::get_scale() const { auto value_ptr = GetAttr(kScale); MS_EXCEPTION_IF_NULL(value_ptr); return GetValue(value_ptr); } -void PowerGrad::set_shift(const float shift) { (void)this->AddAttr(kShift, MakeValue(shift)); } +void PowerGrad::set_shift(const float shift) { (void)this->AddAttr(kShift, api::MakeValue(shift)); } float PowerGrad::get_shift() const { auto value_ptr = GetAttr(kShift); MS_EXCEPTION_IF_NULL(value_ptr); diff --git a/mindspore/core/ops/grad/power_grad.h b/mindspore/core/ops/grad/power_grad.h index 8581203786..c0d68866c0 100644 --- a/mindspore/core/ops/grad/power_grad.h +++ b/mindspore/core/ops/grad/power_grad.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePowerGrad = "PowerGrad"; -class MS_CORE_API PowerGrad : public PrimitiveC { +class MIND_API PowerGrad : public BaseOperator { public: - PowerGrad() : PrimitiveC(kNamePowerGrad) {} - ~PowerGrad() = default; - MS_DECLARE_PARENT(PowerGrad, PrimitiveC); + MIND_API_BASE_MEMBER(PowerGrad); + PowerGrad() : BaseOperator(kNamePowerGrad) {} void Init(const float power, const float scale, const float shift); void set_power(const float power); void set_scale(const float scale); diff --git a/mindspore/core/ops/grad/reciprocal_grad.cc b/mindspore/core/ops/grad/reciprocal_grad.cc index 8efefa8a99..fd5db913d8 100644 --- a/mindspore/core/ops/grad/reciprocal_grad.cc +++ b/mindspore/core/ops/grad/reciprocal_grad.cc @@ -20,6 +20,7 @@ #include "abstract/param_validator.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -48,6 +49,7 @@ TypePtr ReciprocalGradInferType(const PrimitivePtr &primitive, const std::vector } } // namespace +MIND_API_BASE_IMPL(ReciprocalGrad, PrimitiveC, BaseOperator); AbstractBasePtr ReciprocalGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/reciprocal_grad.h b/mindspore/core/ops/grad/reciprocal_grad.h index cdaf6692a7..253f6a66af 100644 --- a/mindspore/core/ops/grad/reciprocal_grad.h +++ b/mindspore/core/ops/grad/reciprocal_grad.h @@ -20,19 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReciprocalGrad = "ReciprocalGrad"; -class ReciprocalGrad : public PrimitiveC { +class MIND_API ReciprocalGrad : public BaseOperator { public: - ReciprocalGrad() : PrimitiveC(kNameReciprocalGrad) {} - ~ReciprocalGrad() = default; - MS_DECLARE_PARENT(ReciprocalGrad, PrimitiveC); + MIND_API_BASE_MEMBER(ReciprocalGrad); + ReciprocalGrad() : BaseOperator(kNameReciprocalGrad) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/relu6_grad.cc b/mindspore/core/ops/grad/relu6_grad.cc index 3a04ce174b..ec9bd18c53 100644 --- a/mindspore/core/ops/grad/relu6_grad.cc +++ b/mindspore/core/ops/grad/relu6_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,6 +61,8 @@ TypePtr ReLU6GradInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/relu6_grad.h b/mindspore/core/ops/grad/relu6_grad.h index 9380f199a1..a3e8406008 100644 --- a/mindspore/core/ops/grad/relu6_grad.h +++ b/mindspore/core/ops/grad/relu6_grad.h @@ -20,18 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReLU6Grad = "ReLU6Grad"; -class ReLU6Grad : public PrimitiveC { +class MIND_API ReLU6Grad : public BaseOperator { public: - ReLU6Grad() : PrimitiveC(kNameReLU6Grad) {} - ~ReLU6Grad() = default; - MS_DECLARE_PARENT(ReLU6Grad, PrimitiveC); + MIND_API_BASE_MEMBER(ReLU6Grad); + ReLU6Grad() : BaseOperator(kNameReLU6Grad) {} void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/relu_grad.cc b/mindspore/core/ops/grad/relu_grad.cc index 7b1cb6f2e5..53e112e186 100644 --- a/mindspore/core/ops/grad/relu_grad.cc +++ b/mindspore/core/ops/grad/relu_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -63,6 +64,8 @@ TypePtr ReLUGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = ReLUGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/relu_grad.h b/mindspore/core/ops/grad/relu_grad.h index d90dd5c98a..216a7a9ae9 100644 --- a/mindspore/core/ops/grad/relu_grad.h +++ b/mindspore/core/ops/grad/relu_grad.h @@ -19,19 +19,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameReLUGrad = prim::kReLUGrad; -class MS_CORE_API ReLUGrad : public PrimitiveC { +constexpr auto kNameReLUGrad = "ReLUGrad"; +class MIND_API ReLUGrad : public BaseOperator { public: - ReLUGrad() : PrimitiveC(prim::kPrimReluGrad->name()) { InitIOName({"x"}, {"output"}); } - ~ReLUGrad() = default; - MS_DECLARE_PARENT(ReLUGrad, PrimitiveC); + MIND_API_BASE_MEMBER(ReLUGrad); + ReLUGrad() : BaseOperator(kNameReLUGrad) { InitIOName({"x"}, {"output"}); } void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/relu_grad_v2.cc b/mindspore/core/ops/grad/relu_grad_v2.cc index 7707c9653a..4a7fca0bba 100644 --- a/mindspore/core/ops/grad/relu_grad_v2.cc +++ b/mindspore/core/ops/grad/relu_grad_v2.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,8 @@ TypePtr ReLUGradV2InferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/relu_grad_v2.h b/mindspore/core/ops/grad/relu_grad_v2.h index 13d1e097d2..f683e6c515 100644 --- a/mindspore/core/ops/grad/relu_grad_v2.h +++ b/mindspore/core/ops/grad/relu_grad_v2.h @@ -19,19 +19,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameReLUGradV2 = prim::kReLUGradV2; -class MS_CORE_API ReLUGradV2 : public PrimitiveC { +constexpr auto kNameReLUGradV2 = "ReLUGradV2"; +class MIND_API ReLUGradV2 : public BaseOperator { public: - ReLUGradV2() : PrimitiveC(prim::kPrimReluGradV2->name()) { InitIOName({"x"}, {"output"}); } - ~ReLUGradV2() = default; - MS_DECLARE_PARENT(ReLUGradV2, PrimitiveC); + MIND_API_BASE_MEMBER(ReLUGradV2); + ReLUGradV2() : BaseOperator(kNameReLUGradV2) { InitIOName({"x"}, {"output"}); } void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/resize_grad.cc b/mindspore/core/ops/grad/resize_grad.cc index d93dcfa9a0..f3f8c2edb8 100644 --- a/mindspore/core/ops/grad/resize_grad.cc +++ b/mindspore/core/ops/grad/resize_grad.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ResizeGrad, PrimitiveC, BaseOperator); void ResizeGrad::Init(const ResizeMethod method, const bool align_corners) { this->set_method(method); this->set_align_corners(align_corners); @@ -32,11 +34,11 @@ void ResizeGrad::Init(const ResizeMethod method, const bool align_corners) { void ResizeGrad::set_method(const ResizeMethod method) { auto swi = (int64_t)method; - (void)this->AddAttr(kMethod, MakeValue(swi)); + (void)this->AddAttr(kMethod, api::MakeValue(swi)); } void ResizeGrad::set_align_corners(const bool align_corners) { - (void)this->AddAttr(kAlignCorners, MakeValue(align_corners)); + (void)this->AddAttr(kAlignCorners, api::MakeValue(align_corners)); } ResizeMethod ResizeGrad::get_method() const { diff --git a/mindspore/core/ops/grad/resize_grad.h b/mindspore/core/ops/grad/resize_grad.h index 0377e4c8cf..b3cbf83523 100644 --- a/mindspore/core/ops/grad/resize_grad.h +++ b/mindspore/core/ops/grad/resize_grad.h @@ -18,18 +18,16 @@ #define MINDSPORE_CORE_OPS_GRAD_RESIZE_GRAD_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameResizeGrad = "ResizeGrad"; -class MS_CORE_API ResizeGrad : public PrimitiveC { +class MIND_API ResizeGrad : public BaseOperator { public: - ResizeGrad() : PrimitiveC(kNameResizeGrad) {} - ~ResizeGrad() = default; - MS_DECLARE_PARENT(ResizeGrad, PrimitiveC); + MIND_API_BASE_MEMBER(ResizeGrad); + ResizeGrad() : BaseOperator(kNameResizeGrad) {} void Init(const ResizeMethod method, const bool align_corners); void set_method(const ResizeMethod method); void set_align_corners(const bool align_corners); @@ -37,8 +35,8 @@ class MS_CORE_API ResizeGrad : public PrimitiveC { bool get_align_corners() const; }; -AbstractBasePtr ResizeGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ResizeGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/resize_nearest_neighbor_grad.cc b/mindspore/core/ops/grad/resize_nearest_neighbor_grad.cc index 7e16e7861d..1288a042e4 100644 --- a/mindspore/core/ops/grad/resize_nearest_neighbor_grad.cc +++ b/mindspore/core/ops/grad/resize_nearest_neighbor_grad.cc @@ -23,6 +23,7 @@ #include "ops/grad/resize_nearest_neighbor_grad.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,6 +61,8 @@ TypePtr ResizeNearestNeighborGradInferType(const PrimitivePtr &prim, const std:: return input_args[0]->BuildType(); } } // namespace + +MIND_API_BASE_IMPL(ResizeNearestNeighborGrad, PrimitiveC, BaseOperator); AbstractBasePtr ResizeNearestNeighborGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto prim_name = primitive->name(); diff --git a/mindspore/core/ops/grad/resize_nearest_neighbor_grad.h b/mindspore/core/ops/grad/resize_nearest_neighbor_grad.h index 128578fdb5..4eca040454 100644 --- a/mindspore/core/ops/grad/resize_nearest_neighbor_grad.h +++ b/mindspore/core/ops/grad/resize_nearest_neighbor_grad.h @@ -20,19 +20,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameResizeNearestNeighborGrad = "ResizeNearestNeighborGrad"; -class ResizeNearestNeighborGrad : public PrimitiveC { +class ResizeNearestNeighborGrad : public BaseOperator { public: - ResizeNearestNeighborGrad() : PrimitiveC(kNameResizeNearestNeighborGrad) {} - ~ResizeNearestNeighborGrad() = default; - MS_DECLARE_PARENT(ResizeNearestNeighborGrad, PrimitiveC); + MIND_API_BASE_MEMBER(ResizeNearestNeighborGrad); + ResizeNearestNeighborGrad() : BaseOperator(kNameResizeNearestNeighborGrad) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/rsqrt_grad.cc b/mindspore/core/ops/grad/rsqrt_grad.cc index 20d1e1655c..aea7e945e1 100644 --- a/mindspore/core/ops/grad/rsqrt_grad.cc +++ b/mindspore/core/ops/grad/rsqrt_grad.cc @@ -15,6 +15,14 @@ */ #include "ops/grad/rsqrt_grad.h" +#include +#include + +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "abstract/param_validator.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +48,7 @@ TypePtr RsqrtGradInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/rsqrt_grad.h b/mindspore/core/ops/grad/rsqrt_grad.h index 7778bd35c8..d2df136337 100644 --- a/mindspore/core/ops/grad/rsqrt_grad.h +++ b/mindspore/core/ops/grad/rsqrt_grad.h @@ -16,28 +16,26 @@ #ifndef MINDSPORE_CORE_OPS_GRAD_RSQRT_GRAD_H_ #define MINDSPORE_CORE_OPS_GRAD_RSQRT_GRAD_H_ + #include #include #include #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRsqrtGrad = "RsqrtGrad"; -class MS_CORE_API RsqrtGrad : public PrimitiveC { +class MIND_API RsqrtGrad : public BaseOperator { public: - RsqrtGrad() : PrimitiveC(kNameRsqrtGrad) { InitIOName({"out_backprop", "input"}, {"output"}); } - ~RsqrtGrad() = default; - MS_DECLARE_PARENT(RsqrtGrad, PrimitiveC); + MIND_API_BASE_MEMBER(RsqrtGrad); + RsqrtGrad() : BaseOperator(kNameRsqrtGrad) { InitIOName({"out_backprop", "input"}, {"output"}); } void Init() const {} }; -AbstractBasePtr RsqrtGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RsqrtGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimRsqrtGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.cc b/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.cc index 8f919ea5d0..f7e577ed2c 100644 --- a/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.cc +++ b/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -64,6 +65,8 @@ TypePtr SigmoidCrossEntropyWithLogitsGradInferType(const PrimitivePtr &primitive return dout_type; } } // namespace + +MIND_API_BASE_IMPL(SigmoidCrossEntropyWithLogitsGrad, PrimitiveC, BaseOperator); AbstractBasePtr SigmoidCrossEntropyWithLogitsGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { diff --git a/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.h b/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.h index 8a6650642d..2fd4ab723f 100644 --- a/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.h +++ b/mindspore/core/ops/grad/sigmoid_cross_entropy_with_logits_grad.h @@ -20,25 +20,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSigmoidCrossEntropyWithLogitsGrad = "SigmoidCrossEntropyWithLogitsGrad"; -class MS_CORE_API SigmoidCrossEntropyWithLogitsGrad : public PrimitiveC { +class MIND_API SigmoidCrossEntropyWithLogitsGrad : public BaseOperator { public: - SigmoidCrossEntropyWithLogitsGrad() : PrimitiveC(kNameSigmoidCrossEntropyWithLogitsGrad) { + MIND_API_BASE_MEMBER(SigmoidCrossEntropyWithLogitsGrad); + SigmoidCrossEntropyWithLogitsGrad() : BaseOperator(kNameSigmoidCrossEntropyWithLogitsGrad) { InitIOName({"x", "y", "dout"}, {"x_grad"}); } - ~SigmoidCrossEntropyWithLogitsGrad() = default; - MS_DECLARE_PARENT(SigmoidCrossEntropyWithLogitsGrad, PrimitiveC); void Init() const {} }; -AbstractBasePtr SigmoidCrossEntropyWithLogitsGradInfer(const abstract::AnalysisEnginePtr &, - const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SigmoidCrossEntropyWithLogitsGradInfer( + const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSigmoidCrossEntropyWithLogitsGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/sigmoid_grad.cc b/mindspore/core/ops/grad/sigmoid_grad.cc index 9e9909aca8..04310a55e2 100644 --- a/mindspore/core/ops/grad/sigmoid_grad.cc +++ b/mindspore/core/ops/grad/sigmoid_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -59,6 +60,8 @@ TypePtr SigmoidGradInfertype(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/sigmoid_grad.h b/mindspore/core/ops/grad/sigmoid_grad.h index 766b70c56c..1eed8f881f 100644 --- a/mindspore/core/ops/grad/sigmoid_grad.h +++ b/mindspore/core/ops/grad/sigmoid_grad.h @@ -20,18 +20,16 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSigmoidGrad = "SigmoidGrad"; -class SigmoidGrad : public PrimitiveC { +class MIND_API SigmoidGrad : public BaseOperator { public: - SigmoidGrad() : PrimitiveC(kNameSigmoidGrad) { InitIOName({"input"}, {"output"}); } - ~SigmoidGrad() = default; - MS_DECLARE_PARENT(SigmoidGrad, PrimitiveC) + MIND_API_BASE_MEMBER(SigmoidGrad); + SigmoidGrad() : BaseOperator(kNameSigmoidGrad) { InitIOName({"input"}, {"output"}); } }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/slice_grad.cc b/mindspore/core/ops/grad/slice_grad.cc index 6208a8274a..87c15f1371 100644 --- a/mindspore/core/ops/grad/slice_grad.cc +++ b/mindspore/core/ops/grad/slice_grad.cc @@ -18,6 +18,7 @@ #include #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -56,6 +57,7 @@ TypePtr SliceGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(SliceGradInferShape(primitive, input_args), SliceGradInferType(primitive, input_args)); diff --git a/mindspore/core/ops/grad/slice_grad.h b/mindspore/core/ops/grad/slice_grad.h index 6959a14ecf..b2c1c3bb82 100644 --- a/mindspore/core/ops/grad/slice_grad.h +++ b/mindspore/core/ops/grad/slice_grad.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSliceGrad = "SliceGrad"; -class MS_CORE_API SliceGrad : public PrimitiveC { +class MIND_API SliceGrad : public BaseOperator { public: - SliceGrad() : PrimitiveC(kNameSliceGrad) { InitIOName({"dy", "x", "begin", "size"}, {"output"}); } - ~SliceGrad() = default; - MS_DECLARE_PARENT(SliceGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SliceGrad); + SliceGrad() : BaseOperator(kNameSliceGrad) { InitIOName({"dy", "x", "begin", "size"}, {"output"}); } void Init() {} }; -AbstractBasePtr SliceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SliceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimSliceGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/smooth_l1_loss_grad.cc b/mindspore/core/ops/grad/smooth_l1_loss_grad.cc index b863b86ed8..317830c07a 100644 --- a/mindspore/core/ops/grad/smooth_l1_loss_grad.cc +++ b/mindspore/core/ops/grad/smooth_l1_loss_grad.cc @@ -21,17 +21,18 @@ #include "ops/grad/smooth_l1_loss_grad.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void SmoothL1LossGrad::Init(const float beta) { this->set_beta(beta); } -void SmoothL1LossGrad::set_beta(const float beta) { (void)this->AddAttr(kBeta, MakeValue(beta)); } +void SmoothL1LossGrad::set_beta(const float beta) { (void)this->AddAttr(kBeta, api::MakeValue(beta)); } float SmoothL1LossGrad::get_beta() const { auto value_ptr = this->GetAttr(kBeta); MS_EXCEPTION_IF_NULL(value_ptr); - return GetValue(value_ptr); + return GetValue(value_ptr); } namespace { @@ -66,6 +67,7 @@ TypePtr SmoothL1LossGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = SmoothL1LossGradInferType(primitive, input_args); diff --git a/mindspore/core/ops/grad/smooth_l1_loss_grad.h b/mindspore/core/ops/grad/smooth_l1_loss_grad.h index 71bdec5fc0..91a39f86d8 100644 --- a/mindspore/core/ops/grad/smooth_l1_loss_grad.h +++ b/mindspore/core/ops/grad/smooth_l1_loss_grad.h @@ -19,25 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSmoothL1LossGrad = "SmoothL1LossGrad"; -class MS_CORE_API SmoothL1LossGrad : public PrimitiveC { +class MIND_API SmoothL1LossGrad : public BaseOperator { public: - SmoothL1LossGrad() : PrimitiveC(kNameSmoothL1LossGrad) {} - ~SmoothL1LossGrad() = default; - MS_DECLARE_PARENT(SmoothL1LossGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SmoothL1LossGrad); + SmoothL1LossGrad() : BaseOperator(kNameSmoothL1LossGrad) {} void Init(); void Init(const float beta); void set_beta(const float beta); float get_beta() const; }; -AbstractBasePtr SmoothL1LossGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SmoothL1LossGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSmoothL1LossGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/soft_margin_loss_grad.cc b/mindspore/core/ops/grad/soft_margin_loss_grad.cc index 11cfc986ac..9c1eab8f01 100644 --- a/mindspore/core/ops/grad/soft_margin_loss_grad.cc +++ b/mindspore/core/ops/grad/soft_margin_loss_grad.cc @@ -19,6 +19,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -54,6 +55,7 @@ TypePtr SoftMarginLossGradInferType(const PrimitivePtr &primitive, const std::ve } } // namespace +MIND_API_BASE_IMPL(SoftMarginLossGrad, PrimitiveC, BaseOperator); AbstractBasePtr SoftMarginLossGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(SoftMarginLossGradInferShape(primitive, input_args), diff --git a/mindspore/core/ops/grad/soft_margin_loss_grad.h b/mindspore/core/ops/grad/soft_margin_loss_grad.h index 98b8bbfdaa..595d022fcf 100644 --- a/mindspore/core/ops/grad/soft_margin_loss_grad.h +++ b/mindspore/core/ops/grad/soft_margin_loss_grad.h @@ -21,22 +21,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftMarginLossGrad = "SoftMarginLossGrad"; -class MS_CORE_API SoftMarginLossGrad : public PrimitiveC { +class MIND_API SoftMarginLossGrad : public BaseOperator { public: - SoftMarginLossGrad() : PrimitiveC(kNameSoftMarginLossGrad) { InitIOName({"predict", "label", "dout"}, {"gradient"}); } - ~SoftMarginLossGrad() = default; - MS_DECLARE_PARENT(SoftMarginLossGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SoftMarginLossGrad); + SoftMarginLossGrad() : BaseOperator(kNameSoftMarginLossGrad) { + InitIOName({"predict", "label", "dout"}, {"gradient"}); + } }; -AbstractBasePtr SoftMarginLossGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SoftMarginLossGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_SOFT_MARGIN_LOSS_GRAD_H_ diff --git a/mindspore/core/ops/grad/soft_shrink_grad.cc b/mindspore/core/ops/grad/soft_shrink_grad.cc index 9f4adfab4a..9bc542376b 100644 --- a/mindspore/core/ops/grad/soft_shrink_grad.cc +++ b/mindspore/core/ops/grad/soft_shrink_grad.cc @@ -25,6 +25,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -54,6 +55,7 @@ TypePtr SoftShrinkGradInferType(const PrimitivePtr &prim, const std::vector &input_args) { return std::make_shared(SoftShrinkGradInferType(primitive, input_args), diff --git a/mindspore/core/ops/grad/soft_shrink_grad.h b/mindspore/core/ops/grad/soft_shrink_grad.h index f6b3b2e587..c034c0b8e2 100644 --- a/mindspore/core/ops/grad/soft_shrink_grad.h +++ b/mindspore/core/ops/grad/soft_shrink_grad.h @@ -21,21 +21,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftShrinkGrad = "SoftShrinkGrad"; -class MS_CORE_API SoftShrinkGrad : public PrimitiveC { +class MIND_API SoftShrinkGrad : public BaseOperator { public: - SoftShrinkGrad() : PrimitiveC(kNameSoftShrinkGrad) { InitIOName({"input_grad", "input_x"}, {"output"}); } - ~SoftShrinkGrad() = default; - MS_DECLARE_PARENT(SoftShrinkGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SoftShrinkGrad); + SoftShrinkGrad() : BaseOperator(kNameSoftShrinkGrad) { InitIOName({"input_grad", "input_x"}, {"output"}); } }; -AbstractBasePtr SoftShrinkGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SoftShrinkGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_SOFTSHRINK_GRAD_H_ diff --git a/mindspore/core/ops/grad/softplus_grad.cc b/mindspore/core/ops/grad/softplus_grad.cc index 7e995d1f39..8fef55c0f7 100644 --- a/mindspore/core/ops/grad/softplus_grad.cc +++ b/mindspore/core/ops/grad/softplus_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -71,6 +72,8 @@ TypePtr SoftplusGradInfertype(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grad/softplus_grad.h b/mindspore/core/ops/grad/softplus_grad.h index a0c8c89f04..60c36c44a8 100644 --- a/mindspore/core/ops/grad/softplus_grad.h +++ b/mindspore/core/ops/grad/softplus_grad.h @@ -13,24 +13,24 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + #ifndef MINDSPORE_CORE_OPS_SOFTPLUS_GRAD_H_ #define MINDSPORE_CORE_OPS_SOFTPLUS_GRAD_H_ #include #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftplusGrad = "SoftplusGrad"; -class SoftplusGrad : public PrimitiveC { +class MIND_API SoftplusGrad : public BaseOperator { public: - SoftplusGrad() : PrimitiveC(kNameSoftplusGrad) { InitIOName({"x"}, {"output"}); } - ~SoftplusGrad() = default; - MS_DECLARE_PARENT(SoftplusGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SoftplusGrad); + SoftplusGrad() : BaseOperator(kNameSoftplusGrad) { InitIOName({"x"}, {"output"}); } }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/sqrt_grad.cc b/mindspore/core/ops/grad/sqrt_grad.cc index 3aefff625e..f409d97f7b 100644 --- a/mindspore/core/ops/grad/sqrt_grad.cc +++ b/mindspore/core/ops/grad/sqrt_grad.cc @@ -18,9 +18,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SqrtGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSqrtGrad, SqrtGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/sqrt_grad.h b/mindspore/core/ops/grad/sqrt_grad.h index 8193747624..20cc080f7f 100644 --- a/mindspore/core/ops/grad/sqrt_grad.h +++ b/mindspore/core/ops/grad/sqrt_grad.h @@ -16,18 +16,16 @@ #ifndef MINDSPORE_CORE_OPS_GRAD_SQRT_GRAD_H_ #define MINDSPORE_CORE_OPS_GRAD_SQRT_GRAD_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSqrtGrad = "SqrtGrad"; -class MS_CORE_API SqrtGrad : public PrimitiveC { +class MIND_API SqrtGrad : public BaseOperator { public: - SqrtGrad() : PrimitiveC(kNameSqrtGrad) { InitIOName({"out_backprop", "input"}, {"output"}); } - ~SqrtGrad() = default; - MS_DECLARE_PARENT(SqrtGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SqrtGrad); + SqrtGrad() : BaseOperator(kNameSqrtGrad) { InitIOName({"out_backprop", "input"}, {"output"}); } void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/grad/strided_slice_grad.cc b/mindspore/core/ops/grad/strided_slice_grad.cc index c9d3976ed4..26b297d1f8 100644 --- a/mindspore/core/ops/grad/strided_slice_grad.cc +++ b/mindspore/core/ops/grad/strided_slice_grad.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -117,6 +118,7 @@ TypePtr StridedSliceGradInferType(const PrimitivePtr &primitive, const std::vect } } // namespace +MIND_API_BASE_IMPL(StridedSliceGrad, PrimitiveC, BaseOperator); AbstractBasePtr StridedSliceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); @@ -129,7 +131,7 @@ AbstractBasePtr StridedSliceGradInfer(const abstract::AnalysisEnginePtr &, const void StridedSliceGrad::set_begin_mask(int64_t begin_mask) { (void)CheckAndConvertUtils::CheckInteger(kBeginMask, begin_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kBeginMask, MakeValue(begin_mask)); + (void)this->AddAttr(kBeginMask, api::MakeValue(begin_mask)); } int64_t StridedSliceGrad::get_begin_mask() const { auto value_ptr = GetAttr(kBeginMask); @@ -138,7 +140,7 @@ int64_t StridedSliceGrad::get_begin_mask() const { } void StridedSliceGrad::set_end_mask(int64_t end_mask) { (void)CheckAndConvertUtils::CheckInteger(kEndMask, end_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kEndMask, MakeValue(end_mask)); + (void)this->AddAttr(kEndMask, api::MakeValue(end_mask)); } int64_t StridedSliceGrad::get_end_mask() const { auto value_ptr = GetAttr(kEndMask); @@ -152,7 +154,7 @@ void StridedSliceGrad::set_ellipsis_mask(int64_t ellipsis_mask) { buffer << "For" << this->name() << ", only support one ellipsis in the index, but got " << this->get_end_mask(); MS_EXCEPTION(ValueError) << buffer.str(); } - (void)this->AddAttr(kEllipsisMask, MakeValue(ellipsis_mask)); + (void)this->AddAttr(kEllipsisMask, api::MakeValue(ellipsis_mask)); } int64_t StridedSliceGrad::get_ellipsis_mask() const { auto value_ptr = GetAttr(kEllipsisMask); @@ -160,7 +162,7 @@ int64_t StridedSliceGrad::get_ellipsis_mask() const { } void StridedSliceGrad::set_new_axis_mask(int64_t new_axis_mask) { (void)CheckAndConvertUtils::CheckInteger(kNewAxisMask, new_axis_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kNewAxisMask, MakeValue(new_axis_mask)); + (void)this->AddAttr(kNewAxisMask, api::MakeValue(new_axis_mask)); } int64_t StridedSliceGrad::get_new_axis_mask() const { auto value_ptr = GetAttr(kNewAxisMask); @@ -168,7 +170,7 @@ int64_t StridedSliceGrad::get_new_axis_mask() const { } void StridedSliceGrad::set_shrink_axis_mask(int64_t shrink_axis_mask) { (void)CheckAndConvertUtils::CheckInteger(kShrinkAxisMask, shrink_axis_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kShrinkAxisMask, MakeValue(shrink_axis_mask)); + (void)this->AddAttr(kShrinkAxisMask, api::MakeValue(shrink_axis_mask)); } int64_t StridedSliceGrad::get_shrink_axis_mask() const { auto value_ptr = GetAttr(kShrinkAxisMask); diff --git a/mindspore/core/ops/grad/strided_slice_grad.h b/mindspore/core/ops/grad/strided_slice_grad.h index bb24de9327..c65cd10d40 100644 --- a/mindspore/core/ops/grad/strided_slice_grad.h +++ b/mindspore/core/ops/grad/strided_slice_grad.h @@ -20,20 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API StridedSliceGrad : public PrimitiveC { +class MIND_API StridedSliceGrad : public BaseOperator { public: - StridedSliceGrad() : PrimitiveC(prim::kPrimStridedSliceGrad->name()) { + MIND_API_BASE_MEMBER(StridedSliceGrad); + StridedSliceGrad() : BaseOperator("StridedSliceGrad") { InitIOName({"dy", "shapex", "begin", "end", "strides"}, {"output"}); } - ~StridedSliceGrad() = default; - MS_DECLARE_PARENT(StridedSliceGrad, PrimitiveC); void Init(int64_t begin_mask = 0, int64_t end_mask = 0, int64_t ellipsis_mask = 0, int64_t new_axis_mask = 0, int64_t shrink_axis_mask = 0); void set_begin_mask(int64_t begin_mask); @@ -48,8 +46,8 @@ class MS_CORE_API StridedSliceGrad : public PrimitiveC { int64_t get_shrink_axis_mask() const; }; -AbstractBasePtr StridedSliceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr StridedSliceGradInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimStridedSliceGradPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/sub_grad.cc b/mindspore/core/ops/grad/sub_grad.cc index 4bf0ca66ec..0d832ab4c2 100644 --- a/mindspore/core/ops/grad/sub_grad.cc +++ b/mindspore/core/ops/grad/sub_grad.cc @@ -17,9 +17,11 @@ #include "ops/grad/sub_grad.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SubGrad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSubGrad, SubGrad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grad/sub_grad.h b/mindspore/core/ops/grad/sub_grad.h index dee3ff3b75..bbd34067bb 100644 --- a/mindspore/core/ops/grad/sub_grad.h +++ b/mindspore/core/ops/grad/sub_grad.h @@ -17,18 +17,16 @@ #define MINDSPORE_CORE_OPS_SUB_GRAD_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSubGrad = "SubGrad"; -class MS_CORE_API SubGrad : public PrimitiveC { +class MIND_API SubGrad : public BaseOperator { public: - SubGrad() : PrimitiveC(kNameSubGrad) {} - ~SubGrad() = default; - MS_DECLARE_PARENT(SubGrad, PrimitiveC); + MIND_API_BASE_MEMBER(SubGrad); + SubGrad() : BaseOperator(kNameSubGrad) {} void Init() const {} }; diff --git a/mindspore/core/ops/grad/tanh_grad.cc b/mindspore/core/ops/grad/tanh_grad.cc index 9a5d754f65..cf60c76625 100644 --- a/mindspore/core/ops/grad/tanh_grad.cc +++ b/mindspore/core/ops/grad/tanh_grad.cc @@ -25,6 +25,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,6 +61,8 @@ TypePtr TanhGradInfertype(const PrimitivePtr &prim, const std::vector &input_args) { auto shape = TanhGradInfershape(primitive, input_args); diff --git a/mindspore/core/ops/grad/tanh_grad.h b/mindspore/core/ops/grad/tanh_grad.h index 64dadee15d..b5b889d2f2 100644 --- a/mindspore/core/ops/grad/tanh_grad.h +++ b/mindspore/core/ops/grad/tanh_grad.h @@ -21,21 +21,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTanhGrad = "TanhGrad"; -class TanhGrad : public PrimitiveC { +class TanhGrad : public BaseOperator { public: + MIND_API_BASE_MEMBER(TanhGrad); /// \brief Constructor. - TanhGrad() : PrimitiveC(kNameTanhGrad) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~TanhGrad() = default; + TanhGrad() : BaseOperator(kNameTanhGrad) { InitIOName({"x"}, {"output"}); } /// \brief Init. - MS_DECLARE_PARENT(TanhGrad, PrimitiveC); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/greater.cc b/mindspore/core/ops/greater.cc index fca9922a6c..ffa53ae5d9 100644 --- a/mindspore/core/ops/greater.cc +++ b/mindspore/core/ops/greater.cc @@ -20,6 +20,8 @@ #include "ops/greater.h" #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,12 +43,14 @@ TypePtr GreaterInferType(const PrimitivePtr &prim, const std::vector(kBool); } } // namespace + +MIND_API_BASE_IMPL(Greater, PrimitiveC, BaseOperator); AbstractBasePtr GreaterInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = GreaterInferType(primitive, input_args); auto infer_shape = GreaterInferShape(primitive, input_args); return abstract::MakeAbstract(infer_shape, infer_type); } -REGISTER_PRIMITIVE_C(kNameGreater, Greater); +REGISTER_PRIMITIVE_EVAL_IMPL(Greater, prim::kPrimGreater, GreaterInfer, nullptr, true); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/greater.h b/mindspore/core/ops/greater.h index 8e5add54f4..8f340ce4bc 100644 --- a/mindspore/core/ops/greater.h +++ b/mindspore/core/ops/greater.h @@ -20,27 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGreater = "Greater"; /// \brief Computes the boolean value of \f$x>y\f$ element-wise. /// Refer to Python API @ref mindspore.ops.Greater for more details. -class MS_CORE_API Greater : public PrimitiveC { +class MIND_API Greater : public BaseOperator { public: + MIND_API_BASE_MEMBER(Greater); /// \brief Constructor. - Greater() : PrimitiveC(kNameGreater) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Greater() = default; - MS_DECLARE_PARENT(Greater, PrimitiveC); + Greater() : BaseOperator(kNameGreater) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Greater for the inputs. void Init() const {} }; -AbstractBasePtr GreaterInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GreaterInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimGreaterPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/greater_equal.cc b/mindspore/core/ops/greater_equal.cc index 325c4aa568..54aacd09d7 100644 --- a/mindspore/core/ops/greater_equal.cc +++ b/mindspore/core/ops/greater_equal.cc @@ -18,6 +18,9 @@ #include #include "ops/greater_equal.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +43,8 @@ TypePtr GreaterEqualInferType(const PrimitivePtr &prim, const std::vector(kBool); } } // namespace + +MIND_API_BASE_IMPL(GreaterEqual, PrimitiveC, BaseOperator); AbstractBasePtr GreaterEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = GreaterEqualInferType(primitive, input_args); diff --git a/mindspore/core/ops/greater_equal.h b/mindspore/core/ops/greater_equal.h index 49d7c4bbbc..52335f2085 100644 --- a/mindspore/core/ops/greater_equal.h +++ b/mindspore/core/ops/greater_equal.h @@ -19,25 +19,21 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameGreaterEqual = "GreaterEqual"; /// \brief Computes the boolean value of \f$x>=y\f$ element-wise. /// Refer to Python API @ref mindspore.ops.GreaterEqual for more details. -class MS_CORE_API GreaterEqual : public PrimitiveC { +class MIND_API GreaterEqual : public BaseOperator { public: + MIND_API_BASE_MEMBER(GreaterEqual); /// \brief Constructor. - GreaterEqual() : PrimitiveC(kNameGreaterEqual) {} - /// \brief Destructor. - ~GreaterEqual() = default; - MS_DECLARE_PARENT(GreaterEqual, PrimitiveC); + GreaterEqual() : BaseOperator(kNameGreaterEqual) {} }; -AbstractBasePtr GreaterEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GreaterEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimGreaterEqualPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/grid_sampler_3d.cc b/mindspore/core/ops/grid_sampler_3d.cc index 3f51a014b8..1d837b9024 100644 --- a/mindspore/core/ops/grid_sampler_3d.cc +++ b/mindspore/core/ops/grid_sampler_3d.cc @@ -16,6 +16,9 @@ #include #include "ops/grid_sampler_3d.h" +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,6 +63,7 @@ TypePtr GridSampler3DInferType(const PrimitivePtr &primitive, const std::vector< } } // namespace +MIND_API_BASE_IMPL(GridSampler3D, PrimitiveC, BaseOperator); AbstractBasePtr GridSampler3DInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/grid_sampler_3d.h b/mindspore/core/ops/grid_sampler_3d.h index 316fac206c..598eb5f30c 100644 --- a/mindspore/core/ops/grid_sampler_3d.h +++ b/mindspore/core/ops/grid_sampler_3d.h @@ -21,22 +21,19 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGridSampler3D = "GridSampler3D"; -class GridSampler3D : public PrimitiveC { +class MIND_API GridSampler3D : public BaseOperator { public: - GridSampler3D() : PrimitiveC(kNameGridSampler3D) { InitIOName({"input_x", "grid"}, {"output"}); } - ~GridSampler3D() = default; - MS_DECLARE_PARENT(GridSampler3D, PrimitiveC); + MIND_API_BASE_MEMBER(GridSampler3D); + GridSampler3D() : BaseOperator(kNameGridSampler3D) { InitIOName({"input_x", "grid"}, {"output"}); } }; -AbstractBasePtr GridSampler3DInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr GridSampler3DInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimGridSampler3D = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/gru.cc b/mindspore/core/ops/gru.cc index ce87310109..3666734d0b 100644 --- a/mindspore/core/ops/gru.cc +++ b/mindspore/core/ops/gru.cc @@ -15,12 +15,17 @@ */ #include "ops/gru.h" +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(GRU, PrimitiveC, BaseOperator); void GRU::Init(bool bidirectional) { this->set_bidirectional(bidirectional); } -void GRU::set_bidirectional(bool bidirectional) { (void)AddAttr(kBidirectional, MakeValue(bidirectional)); } +void GRU::set_bidirectional(bool bidirectional) { (void)AddAttr(kBidirectional, api::MakeValue(bidirectional)); } bool GRU::get_bidirectional() const { auto value_ptr = this->GetAttr(kBidirectional); diff --git a/mindspore/core/ops/gru.h b/mindspore/core/ops/gru.h index 2c31ade354..3e106aa3e3 100644 --- a/mindspore/core/ops/gru.h +++ b/mindspore/core/ops/gru.h @@ -22,29 +22,21 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameGRU = "GRU"; /// \brief GRU defined the GRU operator prototype. -class MS_CORE_API GRU : public PrimitiveC { +class MIND_API GRU : public BaseOperator { public: + MIND_API_BASE_MEMBER(GRU); /// \brief Constructor. - GRU() : PrimitiveC(kNameGRU) { + GRU() : BaseOperator(kNameGRU) { InitIOName({"x", "weight_input", "weight_hidden", "bias_input", "bias_hidden", "seq_length", "init_h"}, {"output", "output_h", "update", "reset", "new", "hidden_new"}); } - - /// \brief Destructor. - ~GRU() = default; - - MS_DECLARE_PARENT(GRU, PrimitiveC); - /// \brief Method to init the op's attributes. /// /// \param[in] bidirectional Define a boolean value to indicate whether the gru is single or double direction. diff --git a/mindspore/core/ops/hashtable_lookup.cc b/mindspore/core/ops/hashtable_lookup.cc index 3bcf696489..bbb08376a8 100644 --- a/mindspore/core/ops/hashtable_lookup.cc +++ b/mindspore/core/ops/hashtable_lookup.cc @@ -19,9 +19,11 @@ #include "utils/check_convert_utils.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(HashtableLookup, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameHashtableLookup, HashtableLookup); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/hashtable_lookup.h b/mindspore/core/ops/hashtable_lookup.h index 2c2dcfa461..f72be909a0 100644 --- a/mindspore/core/ops/hashtable_lookup.h +++ b/mindspore/core/ops/hashtable_lookup.h @@ -18,30 +18,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameHashtableLookup = "HashtableLookup"; /// \brief HashtableLookup defined HashtableLookup operator prototype. -class MS_CORE_API HashtableLookup : public PrimitiveC { +class MIND_API HashtableLookup : public BaseOperator { public: + MIND_API_BASE_MEMBER(HashtableLookup); /// \brief Constructor. - HashtableLookup() : PrimitiveC(kNameHashtableLookup) {} - - /// \brief Destructor. - ~HashtableLookup() = default; - - MS_DECLARE_PARENT(HashtableLookup, PrimitiveC); + HashtableLookup() : BaseOperator(kNameHashtableLookup) {} /// \brief Method to init the op's attributes. void Init() const {} }; -AbstractBasePtr HashtableLookupInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr HashtableLookupInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/hshrink.cc b/mindspore/core/ops/hshrink.cc index cec353a311..a12c6ac944 100644 --- a/mindspore/core/ops/hshrink.cc +++ b/mindspore/core/ops/hshrink.cc @@ -22,6 +22,7 @@ #include "ops/hshrink.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +42,7 @@ TypePtr HShrinkInferType(const PrimitivePtr &primitive, const std::vector &input_args) { return std::make_shared(HShrinkInferType(primitive, input_args), diff --git a/mindspore/core/ops/hshrink.h b/mindspore/core/ops/hshrink.h index 816c9363e9..188c9aeac8 100644 --- a/mindspore/core/ops/hshrink.h +++ b/mindspore/core/ops/hshrink.h @@ -19,26 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameHShrink = "HShrink"; /// \brief Applies the hard shrinkage function element-wise. /// Refer to Python API @ref mindspore.ops.HShrink for more details. -class MS_CORE_API HShrink : public PrimitiveC { +class MIND_API HShrink : public BaseOperator { public: + MIND_API_BASE_MEMBER(HShrink); /// \brief Constructor. - HShrink() : PrimitiveC(kNameHShrink) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~HShrink() = default; - MS_DECLARE_PARENT(HShrink, PrimitiveC); + HShrink() : BaseOperator(kNameHShrink) { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr HShrinkInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr HShrinkInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_HSHRINK_H diff --git a/mindspore/core/ops/hsigmoid.cc b/mindspore/core/ops/hsigmoid.cc index 99ff2535db..d5eb984980 100644 --- a/mindspore/core/ops/hsigmoid.cc +++ b/mindspore/core/ops/hsigmoid.cc @@ -17,6 +17,9 @@ #include #include #include +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +44,8 @@ TypePtr HSigmoidInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(HSigmoid, PrimitiveC, BaseOperator); AbstractBasePtr HSigmoidInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return std::make_shared(HSigmoidInferType(primitive, input_args), diff --git a/mindspore/core/ops/hsigmoid.h b/mindspore/core/ops/hsigmoid.h index 5155d7c1ce..2aca3f0ce1 100644 --- a/mindspore/core/ops/hsigmoid.h +++ b/mindspore/core/ops/hsigmoid.h @@ -19,25 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameHSigmoid = "HSigmoid"; /// \brief Hard sigmoid activation function. Refer to Python API @ref mindspore.ops.HSigmoid for more details. -class MS_CORE_API HSigmoid : public PrimitiveC { +class MIND_API HSigmoid : public BaseOperator { public: + MIND_API_BASE_MEMBER(HSigmoid); /// \brief Constructor. - HSigmoid() : PrimitiveC(kNameHSigmoid) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~HSigmoid() = default; - MS_DECLARE_PARENT(HSigmoid, PrimitiveC); // come from ops/primitive_c.h + HSigmoid() : BaseOperator(kNameHSigmoid) { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr HSigmoidInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr HSigmoidInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/hsv_to_rgb.cc b/mindspore/core/ops/hsv_to_rgb.cc index 88689edf9b..bc799c585e 100644 --- a/mindspore/core/ops/hsv_to_rgb.cc +++ b/mindspore/core/ops/hsv_to_rgb.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,7 @@ TypePtr HSVToRGBInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/hsv_to_rgb.h b/mindspore/core/ops/hsv_to_rgb.h index fecb2f6ffb..cbb9c3fad3 100644 --- a/mindspore/core/ops/hsv_to_rgb.h +++ b/mindspore/core/ops/hsv_to_rgb.h @@ -19,21 +19,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameHSVToRGB = "HSVToRGB"; -class HSVToRGB : public PrimitiveC { +class MIND_API HSVToRGB : public BaseOperator { public: - HSVToRGB() : PrimitiveC(kNameHSVToRGB) { InitIOName({"x"}, {"y"}); } - ~HSVToRGB() = default; - MS_DECLARE_PARENT(HSVToRGB, PrimitiveC); + MIND_API_BASE_MEMBER(HSVToRGB); + HSVToRGB() : BaseOperator(kNameHSVToRGB) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr HSVToRGBInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr HSVToRGBInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using HSVToRGBPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/identity.cc b/mindspore/core/ops/identity.cc index eed84c25d1..71f5e6ac21 100644 --- a/mindspore/core/ops/identity.cc +++ b/mindspore/core/ops/identity.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +42,7 @@ TypePtr IdentityInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/identity.h b/mindspore/core/ops/identity.h index c0902460b7..4119a3752c 100644 --- a/mindspore/core/ops/identity.h +++ b/mindspore/core/ops/identity.h @@ -18,27 +18,24 @@ #define MINDSPORE_CORE_OPS_IDENTITY_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameIdentity = "Identity"; /// \brief Returns a Tensor with the same shape and contents as input. /// Refer to Python API @ref mindspore.ops.Identity for more details. -class MS_CORE_API Identity : public PrimitiveC { +class MIND_API Identity : public BaseOperator { public: + MIND_API_BASE_MEMBER(Identity); /// \brief Constructor. - Identity() : PrimitiveC(kNameIdentity) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Identity() = default; - MS_DECLARE_PARENT(Identity, PrimitiveC); + Identity() : BaseOperator(kNameIdentity) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Identity for the inputs. void Init() const {} }; -AbstractBasePtr IdentityInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr IdentityInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/imag.cc b/mindspore/core/ops/imag.cc index 0f168b29c2..8ac21e0658 100644 --- a/mindspore/core/ops/imag.cc +++ b/mindspore/core/ops/imag.cc @@ -20,6 +20,8 @@ #include #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -58,6 +60,8 @@ AbstractBasePtr ImagInfer(const abstract::AnalysisEnginePtr &, const PrimitivePt return abstract::MakeAbstract(ImagInferShape(primitive, input_args), ImagInferType(primitive, input_args)); } } // namespace + +MIND_API_BASE_IMPL(Imag, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_EVAL_IMPL(Imag, prim::kPrimImag, ImagInfer, nullptr, true); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/imag.h b/mindspore/core/ops/imag.h index 198297cd7e..732d3d8851 100644 --- a/mindspore/core/ops/imag.h +++ b/mindspore/core/ops/imag.h @@ -19,21 +19,18 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns a Tensor that is the imag part of the input. /// Refer to Python API @ref mindspore.ops.Imag for more details. -class MS_CORE_API Imag : public PrimitiveC { +class MIND_API Imag : public BaseOperator { public: + MIND_API_BASE_MEMBER(Imag); /// \brief Constructor. - Imag() : PrimitiveC(prim::kPrimImag->name()) { InitIOName({"input"}, {"output"}); } - /// \brief Destructor. - ~Imag() = default; - MS_DECLARE_PARENT(Imag, PrimitiveC); + Imag() : BaseOperator("Imag") { InitIOName({"input"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Imag for the inputs. void Init() {} }; diff --git a/mindspore/core/ops/index_add.cc b/mindspore/core/ops/index_add.cc index 04a19abd4d..1a37ab825e 100644 --- a/mindspore/core/ops/index_add.cc +++ b/mindspore/core/ops/index_add.cc @@ -18,6 +18,8 @@ #include #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -72,6 +74,7 @@ TypePtr IndexAddInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(IndexAddInferShape(primitive, input_args), IndexAddInferType(primitive, input_args)); diff --git a/mindspore/core/ops/index_add.h b/mindspore/core/ops/index_add.h index 937372c8a3..f9914de983 100644 --- a/mindspore/core/ops/index_add.h +++ b/mindspore/core/ops/index_add.h @@ -21,26 +21,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameIndexAdd = "IndexAdd"; /// \brief Adds tensor y to specified axis and indices of tensor x. /// Refer to Python API @ref mindspore.ops.IndexAdd for more details. -class IndexAdd : public PrimitiveC { +class IndexAdd : public BaseOperator { public: + MIND_API_BASE_MEMBER(IndexAdd); /// \brief Constructor. - IndexAdd() : PrimitiveC(kNameIndexAdd) { InitIOName({"input_x", "indices", "input_y"}, {"output"}); } - /// \brief Destructor. - ~IndexAdd() = default; - MS_DECLARE_PARENT(IndexAdd, PrimitiveC); + IndexAdd() : BaseOperator(kNameIndexAdd) { InitIOName({"input_x", "indices", "input_y"}, {"output"}); } }; -AbstractBasePtr IndexAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr IndexAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/instance_norm.cc b/mindspore/core/ops/instance_norm.cc index bb2bb11116..5caaa5d5b1 100644 --- a/mindspore/core/ops/instance_norm.cc +++ b/mindspore/core/ops/instance_norm.cc @@ -23,12 +23,15 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(InstanceNorm, PrimitiveC, BaseOperator); void InstanceNorm::Init(const float epsilon) { this->set_epsilon(epsilon); } -void InstanceNorm::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, MakeValue(epsilon)); } +void InstanceNorm::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, api::MakeValue(epsilon)); } float InstanceNorm::get_epsilon() const { auto value_ptr = GetAttr(kEpsilon); return GetValue(value_ptr); diff --git a/mindspore/core/ops/instance_norm.h b/mindspore/core/ops/instance_norm.h index 5f55b89d61..ebc630ecc8 100644 --- a/mindspore/core/ops/instance_norm.h +++ b/mindspore/core/ops/instance_norm.h @@ -20,23 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameInstanceNorm = "InstanceNorm"; /// \brief InstanceNorm defined the InstanceNorm operator prototype. -class MS_CORE_API InstanceNorm : public PrimitiveC { +class MIND_API InstanceNorm : public BaseOperator { public: + MIND_API_BASE_MEMBER(InstanceNorm); /// \brief Constructor. - InstanceNorm() : PrimitiveC(kNameInstanceNorm) {} - - /// \brief Destructor. - ~InstanceNorm() = default; - - MS_DECLARE_PARENT(InstanceNorm, PrimitiveC); + InstanceNorm() : BaseOperator(kNameInstanceNorm) {} /// \brief Method to init the op's attributes /// diff --git a/mindspore/core/ops/inv.cc b/mindspore/core/ops/inv.cc index fe8943f74c..50a0340a65 100644 --- a/mindspore/core/ops/inv.cc +++ b/mindspore/core/ops/inv.cc @@ -19,6 +19,9 @@ #include #include #include +#include "ops/primitive_c.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +45,8 @@ TypePtr InvInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/inv.h b/mindspore/core/ops/inv.h index 4f662f4ba5..a58ac2a348 100644 --- a/mindspore/core/ops/inv.h +++ b/mindspore/core/ops/inv.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_INV_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameInv = "Inv"; -class Inv : public PrimitiveC { +class Inv : public BaseOperator { public: - Inv() : PrimitiveC(kNameInv) { InitIOName({"x"}, {"y"}); } - ~Inv() = default; - MS_DECLARE_PARENT(Inv, PrimitiveC); + MIND_API_BASE_MEMBER(Inv); + Inv() : BaseOperator(kNameInv) { InitIOName({"x"}, {"y"}); } void Init() {} }; -AbstractBasePtr InvInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr InvInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimInvPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/invert.cc b/mindspore/core/ops/invert.cc index 20170934e4..0aec3f47fe 100644 --- a/mindspore/core/ops/invert.cc +++ b/mindspore/core/ops/invert.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +41,8 @@ TypePtr InvertInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/invert.h b/mindspore/core/ops/invert.h index ab62bc06ee..8ed6eedf4c 100644 --- a/mindspore/core/ops/invert.h +++ b/mindspore/core/ops/invert.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameInvert = "Invert"; -class Invert : public PrimitiveC { +class MIND_API Invert : public BaseOperator { public: - Invert() : PrimitiveC(kNameInvert) { InitIOName({"x"}, {"y"}); } - ~Invert() = default; - MS_DECLARE_PARENT(Invert, PrimitiveC); + MIND_API_BASE_MEMBER(Invert); + Invert() : BaseOperator(kNameInvert) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr InvertInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr InvertInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimInvertPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/invert_permutation.cc b/mindspore/core/ops/invert_permutation.cc index 6209825210..f1f591f2c5 100644 --- a/mindspore/core/ops/invert_permutation.cc +++ b/mindspore/core/ops/invert_permutation.cc @@ -16,9 +16,12 @@ #include #include "ops/invert_permutation.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(InvertPermutation, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameInvertPermutation, InvertPermutation); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/invert_permutation.h b/mindspore/core/ops/invert_permutation.h index 51cdb8e453..5908374ada 100644 --- a/mindspore/core/ops/invert_permutation.h +++ b/mindspore/core/ops/invert_permutation.h @@ -18,22 +18,19 @@ #define MINDSPORE_CORE_OPS_INVERT_PERMUTATION_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameInvertPermutation = "InvertPermutation"; /// \brief Computes the inverse of an index permutation. /// Refer to Python API @ref mindspore.ops.InvertPermutation for more details. -class MS_CORE_API InvertPermutation : public PrimitiveC { +class MIND_API InvertPermutation : public BaseOperator { public: + MIND_API_BASE_MEMBER(InvertPermutation); /// \brief Constructor. - InvertPermutation() : PrimitiveC(kNameInvertPermutation) {} - /// \brief Destructor. - ~InvertPermutation() = default; - MS_DECLARE_PARENT(InvertPermutation, PrimitiveC); + InvertPermutation() : BaseOperator(kNameInvertPermutation) {} }; } // namespace ops diff --git a/mindspore/core/ops/iou.cc b/mindspore/core/ops/iou.cc index 6fdb6fa9ff..52dccfd553 100644 --- a/mindspore/core/ops/iou.cc +++ b/mindspore/core/ops/iou.cc @@ -17,6 +17,9 @@ #include "ops/iou.h" #include #include +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -75,6 +78,8 @@ TypePtr IOUInferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(IOU, PrimitiveC, BaseOperator); AbstractBasePtr IOUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto type = IOUInferType(primitive, input_args); diff --git a/mindspore/core/ops/iou.h b/mindspore/core/ops/iou.h index 71599ec3ac..431c0f6a5c 100644 --- a/mindspore/core/ops/iou.h +++ b/mindspore/core/ops/iou.h @@ -19,18 +19,15 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class MS_CORE_API IOU : public PrimitiveC { +class MIND_API IOU : public BaseOperator { public: - IOU() : PrimitiveC(prim::kPrimIOU->name()) { InitIOName({"x,y"}, {"output"}); } - ~IOU() = default; - MS_DECLARE_PARENT(IOU, PrimitiveC); + MIND_API_BASE_MEMBER(IOU); + IOU() : BaseOperator("IOU") { InitIOName({"x,y"}, {"output"}); } void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/is_close.cc b/mindspore/core/ops/is_close.cc index 8afb4b7362..0613f3ce7b 100644 --- a/mindspore/core/ops/is_close.cc +++ b/mindspore/core/ops/is_close.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -70,6 +71,8 @@ TypePtr IsCloseInferType(const PrimitivePtr &primitive, const std::vector(kBool); } } // namespace + +MIND_API_BASE_IMPL(IsClose, PrimitiveC, BaseOperator); AbstractBasePtr IsCloseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/is_close.h b/mindspore/core/ops/is_close.h index 8d6189c230..2965bee200 100644 --- a/mindspore/core/ops/is_close.h +++ b/mindspore/core/ops/is_close.h @@ -21,21 +21,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameIsClose = "IsClose"; -class IsClose : public PrimitiveC { +class IsClose : public BaseOperator { public: - IsClose() : PrimitiveC(kNameIsClose) { InitIOName({"x1", "x2"}, {"y"}); } - ~IsClose() = default; - MS_DECLARE_PARENT(IsClose, PrimitiveC); + MIND_API_BASE_MEMBER(IsClose); + IsClose() : BaseOperator(kNameIsClose) { InitIOName({"x1", "x2"}, {"y"}); } }; -AbstractBasePtr IsCloseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr IsCloseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimIsClosePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/is_finite.cc b/mindspore/core/ops/is_finite.cc index 7654b81b10..5692ce4a72 100644 --- a/mindspore/core/ops/is_finite.cc +++ b/mindspore/core/ops/is_finite.cc @@ -16,9 +16,11 @@ #include "ops/is_finite.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(IsFinite, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameIsFinite, IsFinite); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/is_finite.h b/mindspore/core/ops/is_finite.h index 83ec364eb4..6cc57408a9 100644 --- a/mindspore/core/ops/is_finite.h +++ b/mindspore/core/ops/is_finite.h @@ -17,22 +17,19 @@ #ifndef MINDSPORE_CORE_OPS_IS_FINITE_H_ #define MINDSPORE_CORE_OPS_IS_FINITE_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameIsFinite = "IsFinite"; /// \brief Determines which elements are finite for each position. /// Refer to Python API @ref mindspore.ops.IsFinite for more details. -class MS_CORE_API IsFinite : public PrimitiveC { +class MIND_API IsFinite : public BaseOperator { public: + MIND_API_BASE_MEMBER(IsFinite); /// \brief Constructor. - IsFinite() : PrimitiveC(kNameIsFinite) {} - /// \brief Destructor. - ~IsFinite() = default; - MS_DECLARE_PARENT(IsFinite, PrimitiveC); + IsFinite() : BaseOperator(kNameIsFinite) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.IsFinite for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/is_inf.cc b/mindspore/core/ops/is_inf.cc index 2086eb33c0..f070d3eef9 100644 --- a/mindspore/core/ops/is_inf.cc +++ b/mindspore/core/ops/is_inf.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -48,6 +49,7 @@ TypePtr IsInfInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infertype = IsInfInferType(primitive, input_args); diff --git a/mindspore/core/ops/is_inf.h b/mindspore/core/ops/is_inf.h index 2cee02a46c..95dbe95560 100644 --- a/mindspore/core/ops/is_inf.h +++ b/mindspore/core/ops/is_inf.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameIsInf = "IsInf"; -class IsInf : public PrimitiveC { +class MIND_API IsInf : public BaseOperator { public: - IsInf() : PrimitiveC(kNameIsInf) { InitIOName({"x"}, {"y"}); } - ~IsInf() = default; - MS_DECLARE_PARENT(IsInf, PrimitiveC); + MIND_API_BASE_MEMBER(IsInf); + IsInf() : BaseOperator(kNameIsInf) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr IsInfInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr IsInfInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimIsInfPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/is_nan.cc b/mindspore/core/ops/is_nan.cc index f5b49b508f..513dfa8931 100644 --- a/mindspore/core/ops/is_nan.cc +++ b/mindspore/core/ops/is_nan.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -46,6 +47,7 @@ TypePtr IsNanInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/is_nan.h b/mindspore/core/ops/is_nan.h index 46d90bf2c9..693c2d970f 100644 --- a/mindspore/core/ops/is_nan.h +++ b/mindspore/core/ops/is_nan.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameIsNan = "IsNan"; -class IsNan : public PrimitiveC { +class IsNan : public BaseOperator { public: - IsNan() : PrimitiveC(kNameIsNan) { InitIOName({"x"}, {"y"}); } - ~IsNan() = default; - MS_DECLARE_PARENT(IsNan, PrimitiveC); + MIND_API_BASE_MEMBER(IsNan); + IsNan() : BaseOperator(kNameIsNan) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr IsNanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr IsNanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimIsNanPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/l2_loss.cc b/mindspore/core/ops/l2_loss.cc index a961f414e1..5df477f421 100644 --- a/mindspore/core/ops/l2_loss.cc +++ b/mindspore/core/ops/l2_loss.cc @@ -17,6 +17,9 @@ #include "ops/l2_loss.h" #include +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +43,7 @@ TypePtr L2LossInferType(const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = L2LossInferType(primitive, input_args); diff --git a/mindspore/core/ops/l2_loss.h b/mindspore/core/ops/l2_loss.h index 82daa14292..07c711eae0 100644 --- a/mindspore/core/ops/l2_loss.h +++ b/mindspore/core/ops/l2_loss.h @@ -19,23 +19,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameL2Loss = "L2Loss"; -class MS_CORE_API L2Loss : public PrimitiveC { +class MIND_API L2Loss : public BaseOperator { public: - L2Loss() : PrimitiveC(kNameL2Loss) { InitIOName({"x"}, {"output"}); } - ~L2Loss() = default; - MS_DECLARE_PARENT(L2Loss, PrimitiveC); + MIND_API_BASE_MEMBER(L2Loss); + L2Loss() : BaseOperator(kNameL2Loss) { InitIOName({"x"}, {"output"}); } void Init() {} }; -AbstractBasePtr L2LossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr L2LossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimL2LossPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/l2_normalize.cc b/mindspore/core/ops/l2_normalize.cc index 7e8d757a5c..90a61a0cf4 100644 --- a/mindspore/core/ops/l2_normalize.cc +++ b/mindspore/core/ops/l2_normalize.cc @@ -17,17 +17,21 @@ #include #include "ops/l2_normalize.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(L2Normalize, PrimitiveC, BaseOperator); void L2Normalize::Init(const std::vector &axis, const float epsilon) { this->set_axis(axis); this->set_epsilon(epsilon); } -void L2Normalize::set_axis(const std::vector &axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void L2Normalize::set_axis(const std::vector &axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } -void L2Normalize::set_epsilon(const float epsilon) { (void)AddAttr(kEpsilon, MakeValue(epsilon)); } +void L2Normalize::set_epsilon(const float epsilon) { (void)AddAttr(kEpsilon, api::MakeValue(epsilon)); } std::vector L2Normalize::get_axis() const { return GetValue>(GetAttr(kAxis)); } diff --git a/mindspore/core/ops/l2_normalize.h b/mindspore/core/ops/l2_normalize.h index 1ce0322f44..be993ba351 100644 --- a/mindspore/core/ops/l2_normalize.h +++ b/mindspore/core/ops/l2_normalize.h @@ -19,22 +19,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameL2Normalize = "L2Normalize"; /// \brief L2 Normalization Operator. Refer to Python API @ref mindspore.ops.L2Normalize for more details. -class MS_CORE_API L2Normalize : public PrimitiveC { +class MIND_API L2Normalize : public BaseOperator { public: + MIND_API_BASE_MEMBER(L2Normalize); /// \brief Constructor. - explicit L2Normalize(const std::string &name = kNameL2Normalize) : PrimitiveC(name) {} - /// \brief Destructor. - ~L2Normalize() = default; - MS_DECLARE_PARENT(L2Normalize, PrimitiveC); + explicit L2Normalize(const std::string &name = kNameL2Normalize) : BaseOperator(name) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.L2Normalize for the inputs. void Init(const std::vector &axis, const float epsilon = 1e-4); /// \brief Set axis. @@ -50,8 +46,8 @@ class MS_CORE_API L2Normalize : public PrimitiveC { /// \return epsilon. float get_epsilon() const; }; -AbstractBasePtr L2NormalizeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr L2NormalizeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_L2_NORMALIZE_H_ diff --git a/mindspore/core/ops/lars_v2_update.cc b/mindspore/core/ops/lars_v2_update.cc index 376beb1f1f..57dfa8e613 100644 --- a/mindspore/core/ops/lars_v2_update.cc +++ b/mindspore/core/ops/lars_v2_update.cc @@ -17,6 +17,7 @@ #include "ops/lars_v2_update.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -90,6 +91,8 @@ TypePtr LARSUpdateInferType(const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = LARSUpdateInferType(primitive, input_args); diff --git a/mindspore/core/ops/lars_v2_update.h b/mindspore/core/ops/lars_v2_update.h index 88077a0c5f..85f3871757 100644 --- a/mindspore/core/ops/lars_v2_update.h +++ b/mindspore/core/ops/lars_v2_update.h @@ -21,22 +21,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLARSUpdate = "LARSUpdate"; -class MS_CORE_API LARSUpdate : public PrimitiveC { +class MIND_API LARSUpdate : public BaseOperator { public: - explicit LARSUpdate(const std::string &name = kNameLARSUpdate) : PrimitiveC(name) {} - ~LARSUpdate() = default; - MS_DECLARE_PARENT(LARSUpdate, PrimitiveC); + MIND_API_BASE_MEMBER(LARSUpdate); + explicit LARSUpdate(const std::string &name = kNameLARSUpdate) : BaseOperator(name) {} }; -AbstractBasePtr LARSUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LARSUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLARSUpdatePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/layer_norm.cc b/mindspore/core/ops/layer_norm.cc index aa5276a520..c4122e4973 100644 --- a/mindspore/core/ops/layer_norm.cc +++ b/mindspore/core/ops/layer_norm.cc @@ -17,6 +17,7 @@ #include "ops/layer_norm.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -35,6 +36,7 @@ ShapeVector CalLayerNormMeanAndVarShape(int64_t begin_norm_axis, const ShapeVect } } // namespace +MIND_API_BASE_IMPL(LayerNorm, PrimitiveC, BaseOperator); AbstractBasePtr LayerNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { // Inputs: three tensors(x, gamma, beta). @@ -129,12 +131,12 @@ void LayerNorm::Init(const int64_t begin_norm_axis, const int64_t begin_params_a this->set_epsilon(epsilon); } void LayerNorm::set_begin_norm_axis(const int64_t begin_norm_axis) { - (void)this->AddAttr(kBeginNormAxis, MakeValue(begin_norm_axis)); + (void)this->AddAttr(kBeginNormAxis, api::MakeValue(begin_norm_axis)); } void LayerNorm::set_begin_params_axis(const int64_t begin_params_axis) { - (void)this->AddAttr(kBeginParamsAxis, MakeValue(begin_params_axis)); + (void)this->AddAttr(kBeginParamsAxis, api::MakeValue(begin_params_axis)); } -void LayerNorm::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, MakeValue(epsilon)); } +void LayerNorm::set_epsilon(const float epsilon) { (void)this->AddAttr(kEpsilon, api::MakeValue(epsilon)); } int64_t LayerNorm::get_begin_norm_axis() const { auto value_ptr = this->GetAttr(kBeginNormAxis); diff --git a/mindspore/core/ops/layer_norm.h b/mindspore/core/ops/layer_norm.h index 56a7c35681..9356e1611f 100644 --- a/mindspore/core/ops/layer_norm.h +++ b/mindspore/core/ops/layer_norm.h @@ -20,23 +20,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameLayerNorm = prim::kLayerNorm; +constexpr auto kNameLayerNorm = "LayerNorm"; /// \brief Applies the Layer Normalization to the input tensor. /// Refer to Python API @ref mindspore.ops.LayerNorm for more details. -class MS_CORE_API LayerNorm : public PrimitiveC { +class MIND_API LayerNorm : public BaseOperator { public: + MIND_API_BASE_MEMBER(LayerNorm); /// \brief Constructor. - LayerNorm() : PrimitiveC(kNameLayerNorm) {} - explicit LayerNorm(const std::string k_name) : PrimitiveC(k_name) {} - /// \brief Destructor. - ~LayerNorm() = default; - MS_DECLARE_PARENT(LayerNorm, PrimitiveC); + LayerNorm() : BaseOperator(kNameLayerNorm) {} + explicit LayerNorm(const std::string k_name) : BaseOperator(k_name) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LayerNorm for the inputs. void Init(const int64_t begin_norm_axis = 1, const int64_t begin_params_axis = 1, const float epsilon = 1e-7); /// \brief Set begin_norm_axis. @@ -59,8 +56,8 @@ class MS_CORE_API LayerNorm : public PrimitiveC { float get_epsilon() const; }; -AbstractBasePtr LayerNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLayerNormPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.cc b/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.cc index 22e7849feb..0aa0c2c32a 100644 --- a/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.cc +++ b/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,10 +46,12 @@ TypePtr LayerNormBetaGammaBackpropV2InferType(const PrimitivePtr &prim, } } // namespace +MIND_API_BASE_IMPL(LayerNormBetaGammaBackpropV2, PrimitiveC, BaseOperator); void LayerNormBetaGammaBackpropV2::Init(const std::vector &shape_gamma) { set_shape_gamma(shape_gamma); } void LayerNormBetaGammaBackpropV2::set_shape_gamma(const std::vector &shape_gamma) { - (void)AddAttr(kShapeGamma, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kShapeGamma, shape_gamma, name()))); + (void)AddAttr(kShapeGamma, + api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kShapeGamma, shape_gamma, name()))); } std::vector LayerNormBetaGammaBackpropV2::get_shape_gamma() const { diff --git a/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.h b/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.h index a45d113d29..7a5737b21c 100644 --- a/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.h +++ b/mindspore/core/ops/layer_norm_beta_gamma_backprop_v2.h @@ -20,25 +20,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class LayerNormBetaGammaBackpropV2 : public PrimitiveC { +class LayerNormBetaGammaBackpropV2 : public BaseOperator { public: - LayerNormBetaGammaBackpropV2() : PrimitiveC(prim::kPrimLayerNormBetaGammaBackpropV2->name()) {} - ~LayerNormBetaGammaBackpropV2() = default; - MS_DECLARE_PARENT(LayerNormBetaGammaBackpropV2, PrimitiveC); + MIND_API_BASE_MEMBER(LayerNormBetaGammaBackpropV2); + LayerNormBetaGammaBackpropV2() : BaseOperator("LayerNormBetaGammaBackpropV2") {} void Init(const std::vector &shape_gamma); void set_shape_gamma(const std::vector &shape_gamma); std::vector get_shape_gamma() const; }; -AbstractBasePtr LayerNormBetaGammaBackpropV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormBetaGammaBackpropV2Infer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/layer_norm_x_backprop_v2.cc b/mindspore/core/ops/layer_norm_x_backprop_v2.cc index 1901474dc0..58449d4c04 100644 --- a/mindspore/core/ops/layer_norm_x_backprop_v2.cc +++ b/mindspore/core/ops/layer_norm_x_backprop_v2.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,6 +45,7 @@ TypePtr LayerNormXBackpropV2InferType(const PrimitivePtr &prim, const std::vecto } } // namespace +MIND_API_BASE_IMPL(LayerNormXBackpropV2, PrimitiveC, BaseOperator); AbstractBasePtr LayerNormXBackpropV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/layer_norm_x_backprop_v2.h b/mindspore/core/ops/layer_norm_x_backprop_v2.h index 9f517d01a1..0a047b435f 100644 --- a/mindspore/core/ops/layer_norm_x_backprop_v2.h +++ b/mindspore/core/ops/layer_norm_x_backprop_v2.h @@ -20,23 +20,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -class LayerNormXBackpropV2 : public PrimitiveC { +constexpr auto kNameLayerNormXBackpropV2 = "LayerNormXBackpropV2"; +class MIND_API LayerNormXBackpropV2 : public BaseOperator { public: - LayerNormXBackpropV2() : PrimitiveC(prim::kPrimLayerNormXBackpropV2->name()) {} - ~LayerNormXBackpropV2() = default; - MS_DECLARE_PARENT(LayerNormXBackpropV2, PrimitiveC); + MIND_API_BASE_MEMBER(LayerNormXBackpropV2); + LayerNormXBackpropV2() : BaseOperator(kNameLayerNormXBackpropV2) {} void Init() const {} }; -AbstractBasePtr LayerNormXBackpropV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LayerNormXBackpropV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/leaky_relu.cc b/mindspore/core/ops/leaky_relu.cc index 90f895f090..e4341b48fd 100644 --- a/mindspore/core/ops/leaky_relu.cc +++ b/mindspore/core/ops/leaky_relu.cc @@ -15,16 +15,20 @@ */ #include "ops/leaky_relu.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void LeakyRelu::Init(const float negative_slope) { this->set_negative_slope(negative_slope); } void LeakyRelu::set_negative_slope(const float negative_slope) { - (void)this->AddAttr(kNegativeSlope, MakeValue(negative_slope)); + (void)this->AddAttr(kNegativeSlope, api::MakeValue(negative_slope)); } float LeakyRelu::get_negative_slope() const { return GetValue(GetAttr(kNegativeSlope)); } +MIND_API_BASE_IMPL(LeakyRelu, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLeakyRelu, LeakyRelu); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/leaky_relu.h b/mindspore/core/ops/leaky_relu.h index 5c14b3a502..0c3645f6bb 100644 --- a/mindspore/core/ops/leaky_relu.h +++ b/mindspore/core/ops/leaky_relu.h @@ -20,22 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLeakyRelu = "LeakyRelu"; /// \brief Leaky ReLU activation function. Refer to Python API @ref mindspore.nn.LeakyReLU for more details. -class MS_CORE_API LeakyRelu : public PrimitiveC { +class MIND_API LeakyRelu : public BaseOperator { public: + MIND_API_BASE_MEMBER(LeakyRelu); /// \brief Constructor. - LeakyRelu() : PrimitiveC(kNameLeakyRelu) {} - /// \brief Destructor. - ~LeakyRelu() = default; - MS_DECLARE_PARENT(LeakyRelu, PrimitiveC); + LeakyRelu() : BaseOperator(kNameLeakyRelu) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.nn.LeakyReLU for the inputs. void Init(const float negative_slope); /// \brief Set negative_slope. @@ -46,8 +42,8 @@ class MS_CORE_API LeakyRelu : public PrimitiveC { float get_negative_slope() const; }; -AbstractBasePtr LeakyReluInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LeakyReluInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lerp.cc b/mindspore/core/ops/lerp.cc index 6292e3fc01..b7441aa695 100644 --- a/mindspore/core/ops/lerp.cc +++ b/mindspore/core/ops/lerp.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -67,6 +68,7 @@ TypePtr LerpInferType(const PrimitivePtr &prim, const std::vector &input_args) { return std::make_shared(LerpInferType(primitive, input_args), diff --git a/mindspore/core/ops/lerp.h b/mindspore/core/ops/lerp.h index aa6def9c90..a2fe905cb9 100644 --- a/mindspore/core/ops/lerp.h +++ b/mindspore/core/ops/lerp.h @@ -19,27 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLerp = "Lerp"; /// \brief Does a linear interpolation of two tensors start and end based on a float or tensor weight. /// Refer to Python API @ref mindspore.ops.Lerp for more details. -class Lerp : public PrimitiveC { +class MIND_API Lerp : public BaseOperator { public: + MIND_API_BASE_MEMBER(Lerp); /// \brief Constructor. - Lerp() : PrimitiveC(kNameLerp) { InitIOName({"start", "end", "weight"}, {"output"}); } - /// \brief Destructor. - ~Lerp() = default; - MS_DECLARE_PARENT(Lerp, PrimitiveC); + Lerp() : BaseOperator(kNameLerp) { InitIOName({"start", "end", "weight"}, {"output"}); } }; -AbstractBasePtr LerpInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LerpInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_LERP_H_ diff --git a/mindspore/core/ops/less.cc b/mindspore/core/ops/less.cc index a8278e60f6..1a8bed003f 100644 --- a/mindspore/core/ops/less.cc +++ b/mindspore/core/ops/less.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -43,6 +44,7 @@ TypePtr LessInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto shape = LessInferShape(primitive, input_args); diff --git a/mindspore/core/ops/less.h b/mindspore/core/ops/less.h index 2896d9ef2f..7914fa8829 100644 --- a/mindspore/core/ops/less.h +++ b/mindspore/core/ops/less.h @@ -19,26 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLess = "Less"; /// \brief Computes the boolean value of \f$x &input_args); +abstract::AbstractBasePtr LessInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_LESS_H_ diff --git a/mindspore/core/ops/less_equal.cc b/mindspore/core/ops/less_equal.cc index 19d303bbc9..44935bbefe 100644 --- a/mindspore/core/ops/less_equal.cc +++ b/mindspore/core/ops/less_equal.cc @@ -20,6 +20,8 @@ #include "ops/less_equal.h" #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +44,7 @@ TypePtr LessEqualInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_shape = LessEqualInferShape(primitive, input_args); diff --git a/mindspore/core/ops/less_equal.h b/mindspore/core/ops/less_equal.h index 5c158b3c7f..c3857be0b7 100644 --- a/mindspore/core/ops/less_equal.h +++ b/mindspore/core/ops/less_equal.h @@ -19,28 +19,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLessEqual = "LessEqual"; /// \brief Computes the boolean value of \f$x<=y\f$ element-wise. /// Refer to Python API @ref mindspore.ops.LessEqual for more details. -class MS_CORE_API LessEqual : public PrimitiveC { +class MIND_API LessEqual : public BaseOperator { public: + MIND_API_BASE_MEMBER(LessEqual); /// \brief Constructor. - LessEqual() : PrimitiveC(kNameLessEqual) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~LessEqual() = default; - MS_DECLARE_PARENT(LessEqual, PrimitiveC); + LessEqual() : BaseOperator(kNameLessEqual) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LessEqual for the inputs. void Init() const {} }; -AbstractBasePtr LessEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LessEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lin_space.cc b/mindspore/core/ops/lin_space.cc index 82159a49ae..089a232f40 100644 --- a/mindspore/core/ops/lin_space.cc +++ b/mindspore/core/ops/lin_space.cc @@ -16,9 +16,12 @@ #include "ops/lin_space.h" #include +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(LinSpace, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLinSpace, LinSpace); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lin_space.h b/mindspore/core/ops/lin_space.h index 8c68c18b62..2715e469c2 100644 --- a/mindspore/core/ops/lin_space.h +++ b/mindspore/core/ops/lin_space.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLinSpace = "LinSpace"; /// \brief Returns a Tensor whose value is evenly spaced in the interval start and stop (including start and stop). /// Refer to Python API @ref mindspore.ops.LinSpace for more details. -class MS_CORE_API LinSpace : public PrimitiveC { +class MIND_API LinSpace : public BaseOperator { public: + MIND_API_BASE_MEMBER(LinSpace); /// \brief Constructor. - LinSpace() : PrimitiveC(kNameLinSpace) { InitIOName({"start", "stop", "num"}, {"output"}); } - /// \brief Destructor. - ~LinSpace() = default; - MS_DECLARE_PARENT(LinSpace, PrimitiveC); + LinSpace() : BaseOperator(kNameLinSpace) { InitIOName({"start", "stop", "num"}, {"output"}); } }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/log.cc b/mindspore/core/ops/log.cc index e48f018fcc..64992c53d8 100644 --- a/mindspore/core/ops/log.cc +++ b/mindspore/core/ops/log.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -59,6 +60,7 @@ TypePtr LogInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(LogInferShape(primitive, input_args), LogInferType(primitive, input_args)); diff --git a/mindspore/core/ops/log.h b/mindspore/core/ops/log.h index ecf0ff56d2..b59159ebbb 100644 --- a/mindspore/core/ops/log.h +++ b/mindspore/core/ops/log.h @@ -19,25 +19,21 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { -constexpr auto kNameLog = prim::kLog; +constexpr auto kNameLog = "Log"; /// \brief Returns the natural logarithm of a tensor element-wise. /// Refer to Python API @ref mindspore.ops.Log for more details. -class MS_CORE_API Log : public PrimitiveC { +class MIND_API Log : public BaseOperator { public: + MIND_API_BASE_MEMBER(Log); /// \brief Constructor. - Log() : PrimitiveC(prim::kPrimLog->name()) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Log() = default; - MS_DECLARE_PARENT(Log, PrimitiveC); + Log() : BaseOperator("Log") { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr LogInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_LOG_H_ diff --git a/mindspore/core/ops/log1p.cc b/mindspore/core/ops/log1p.cc index 709214f9ab..cd3b1690e7 100644 --- a/mindspore/core/ops/log1p.cc +++ b/mindspore/core/ops/log1p.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -46,6 +47,7 @@ TypePtr Log1pInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/log1p.h b/mindspore/core/ops/log1p.h index f180a628a2..bf5756ee0b 100644 --- a/mindspore/core/ops/log1p.h +++ b/mindspore/core/ops/log1p.h @@ -20,22 +20,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns the natural logarithm of one plus the input tensor element-wise. /// Refer to Python API @ref mindspore.ops.Log1p for more details. -class MS_CORE_API Log1p : public PrimitiveC { +class MIND_API Log1p : public BaseOperator { public: + MIND_API_BASE_MEMBER(Log1p); /// \brief Constructor. - Log1p() : PrimitiveC(prim::kPrimLog1p->name()) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Log1p() = default; - MS_DECLARE_PARENT(Log1p, PrimitiveC); + Log1p() : BaseOperator("Log1p") { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Log1p for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/log_matrix_determinant.cc b/mindspore/core/ops/log_matrix_determinant.cc index b444f6765c..bd2511d2a6 100644 --- a/mindspore/core/ops/log_matrix_determinant.cc +++ b/mindspore/core/ops/log_matrix_determinant.cc @@ -20,6 +20,8 @@ #include "ops/op_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -48,6 +50,7 @@ TuplePtr LogMatrixDeterminantInferType(const PrimitivePtr &prim, const std::vect } } // namespace +MIND_API_BASE_IMPL(LogMatrixDeterminant, PrimitiveC, BaseOperator); AbstractBasePtr LogMatrixDeterminantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/log_matrix_determinant.h b/mindspore/core/ops/log_matrix_determinant.h index 961ea46fd4..856a1ee650 100644 --- a/mindspore/core/ops/log_matrix_determinant.h +++ b/mindspore/core/ops/log_matrix_determinant.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogMatrixDeterminant = "LogMatrixDeterminant"; -class LogMatrixDeterminant : public PrimitiveC { +class MIND_API LogMatrixDeterminant : public BaseOperator { public: - LogMatrixDeterminant() : PrimitiveC(kNameLogMatrixDeterminant) { InitIOName({"x"}, {"sign", "output"}); } - ~LogMatrixDeterminant() = default; - MS_DECLARE_PARENT(LogMatrixDeterminant, PrimitiveC); + MIND_API_BASE_MEMBER(LogMatrixDeterminant); + LogMatrixDeterminant() : BaseOperator(kNameLogMatrixDeterminant) { InitIOName({"x"}, {"sign", "output"}); } }; -AbstractBasePtr LogMatrixDeterminantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogMatrixDeterminantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLogMatrixDeterminantPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/log_softmax.cc b/mindspore/core/ops/log_softmax.cc index 809b6c04e9..74fb9f0a7f 100644 --- a/mindspore/core/ops/log_softmax.cc +++ b/mindspore/core/ops/log_softmax.cc @@ -23,10 +23,12 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void LogSoftmax::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +MIND_API_BASE_IMPL(LogSoftmax, PrimitiveC, BaseOperator); +void LogSoftmax::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t LogSoftmax::get_axis() const { return GetValue(GetAttr(kAxis)); } diff --git a/mindspore/core/ops/log_softmax.h b/mindspore/core/ops/log_softmax.h index 4c502c5d73..607b205d06 100644 --- a/mindspore/core/ops/log_softmax.h +++ b/mindspore/core/ops/log_softmax.h @@ -21,21 +21,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogSoftmax = "LogSoftmax"; /// \brief Log Softmax activation function. Refer to Python API @ref mindspore.ops.LogSoftmax for more details. -class MS_CORE_API LogSoftmax : public PrimitiveC { +class MIND_API LogSoftmax : public BaseOperator { public: + MIND_API_BASE_MEMBER(LogSoftmax); /// \brief Constructor. - LogSoftmax() : PrimitiveC(kNameLogSoftmax) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~LogSoftmax() = default; - MS_DECLARE_PARENT(LogSoftmax, PrimitiveC); + LogSoftmax() : BaseOperator(kNameLogSoftmax) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LogSoftmax for the inputs. void Init(const int64_t axis = -1); /// \brief Set axis. diff --git a/mindspore/core/ops/logical_and.cc b/mindspore/core/ops/logical_and.cc index 3e591d4337..2e42f3b37b 100644 --- a/mindspore/core/ops/logical_and.cc +++ b/mindspore/core/ops/logical_and.cc @@ -23,6 +23,7 @@ #include "ops/logical_and.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,6 +45,7 @@ TypePtr LogicalAndInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/logical_and.h b/mindspore/core/ops/logical_and.h index 0ba6da0b19..0f7ebd3ef4 100644 --- a/mindspore/core/ops/logical_and.h +++ b/mindspore/core/ops/logical_and.h @@ -20,27 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogicalAnd = "LogicalAnd"; /// \brief Computes the "logical AND" of two tensors element-wise. /// Refer to Python API @ref mindspore.ops.LogicalAnd for more details. -class MS_CORE_API LogicalAnd : public PrimitiveC { +class MIND_API LogicalAnd : public BaseOperator { public: + MIND_API_BASE_MEMBER(LogicalAnd); /// \brief Constructor. - LogicalAnd() : PrimitiveC(kNameLogicalAnd) { InitIOName({"x1", "x2"}, {"y"}); } - /// \brief Destructor. - ~LogicalAnd() = default; - MS_DECLARE_PARENT(LogicalAnd, PrimitiveC); + LogicalAnd() : BaseOperator(kNameLogicalAnd) { InitIOName({"x1", "x2"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LogicalAnd for the inputs. void Init() const {} }; -AbstractBasePtr LogicalAndInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogicalAndInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimLogicalAndPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/logical_not.cc b/mindspore/core/ops/logical_not.cc index e6edb1bea2..63ab2c8ea3 100644 --- a/mindspore/core/ops/logical_not.cc +++ b/mindspore/core/ops/logical_not.cc @@ -23,6 +23,7 @@ #include "ops/logical_not.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -47,6 +48,7 @@ TypePtr LogicalNotInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/logical_not.h b/mindspore/core/ops/logical_not.h index fd03f79da7..104703b959 100644 --- a/mindspore/core/ops/logical_not.h +++ b/mindspore/core/ops/logical_not.h @@ -18,27 +18,24 @@ #define MINDSPORE_CORE_OPS_LOGICAL_NOT_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogicalNot = "LogicalNot"; /// \brief Computes the "logical NOT" of a tensor element-wise. /// Refer to Python API @ref mindspore.ops.LogicalNot for more details. -class MS_CORE_API LogicalNot : public PrimitiveC { +class MIND_API LogicalNot : public BaseOperator { public: + MIND_API_BASE_MEMBER(LogicalNot); /// \brief Constructor. - LogicalNot() : PrimitiveC(kNameLogicalNot) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~LogicalNot() = default; - MS_DECLARE_PARENT(LogicalNot, PrimitiveC); + LogicalNot() : BaseOperator(kNameLogicalNot) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LogicalNot for the inputs. void Init() const {} }; -AbstractBasePtr LogicalNotInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogicalNotInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimLogicalNotPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/logical_or.cc b/mindspore/core/ops/logical_or.cc index e4e0bf4cff..a581dfadc1 100644 --- a/mindspore/core/ops/logical_or.cc +++ b/mindspore/core/ops/logical_or.cc @@ -23,6 +23,7 @@ #include "ops/logical_or.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,6 +45,7 @@ TypePtr LogicalOrInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/logical_or.h b/mindspore/core/ops/logical_or.h index ffee0f4222..c5cfa17953 100644 --- a/mindspore/core/ops/logical_or.h +++ b/mindspore/core/ops/logical_or.h @@ -20,27 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogicalOr = "LogicalOr"; /// \brief Computes the "logical OR" of two tensors element-wise. /// Refer to Python API @ref mindspore.ops.LogicalOr for more details. -class MS_CORE_API LogicalOr : public PrimitiveC { +class MIND_API LogicalOr : public BaseOperator { public: + MIND_API_BASE_MEMBER(LogicalOr); /// \brief Constructor. - LogicalOr() : PrimitiveC(kNameLogicalOr) { InitIOName({"x1", "x2"}, {"y"}); } - /// \brief Destructor. - ~LogicalOr() = default; - MS_DECLARE_PARENT(LogicalOr, PrimitiveC); + LogicalOr() : BaseOperator(kNameLogicalOr) { InitIOName({"x1", "x2"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LogicalOr for the inputs. void Init() const {} }; -AbstractBasePtr LogicalOrInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogicalOrInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimLogicalOrPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/logical_xor.cc b/mindspore/core/ops/logical_xor.cc index eae6031028..f7d543962f 100644 --- a/mindspore/core/ops/logical_xor.cc +++ b/mindspore/core/ops/logical_xor.cc @@ -14,11 +14,14 @@ * limitations under the License. */ +#include "ops/logical_xor.h" #include #include #include #include "ops/op_utils.h" -#include "ops/logical_xor.h" +#include "base/core_ops.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -38,6 +41,7 @@ TypePtr LogicalXorInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/logical_xor.h b/mindspore/core/ops/logical_xor.h index b140679753..328010f8a7 100644 --- a/mindspore/core/ops/logical_xor.h +++ b/mindspore/core/ops/logical_xor.h @@ -18,28 +18,26 @@ #define MINDSPORE_CORE_OPS_LOGICAL_XOR_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLogicalXor = "LogicalXor"; /// \brief Computes the truth value of x1 XOR x2, element-wise. /// Refer to Python API @ref mindspore.numpy.logical_xor for more details. -class MS_CORE_API LogicalXor : public PrimitiveC { +class MIND_API LogicalXor : public BaseOperator { public: + MIND_API_BASE_MEMBER(LogicalXor); /// \brief Constructor. - LogicalXor() : PrimitiveC(kNameLogicalXor) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~LogicalXor() = default; - MS_DECLARE_PARENT(LogicalXor, PrimitiveC); + LogicalXor() : BaseOperator(kNameLogicalXor) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.numpy.logical_xor for the inputs. void Init() const {} }; -AbstractBasePtr LogicalXorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LogicalXorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLogicalXorPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lower_bound.cc b/mindspore/core/ops/lower_bound.cc index a538ac379f..5eb0be26e1 100644 --- a/mindspore/core/ops/lower_bound.cc +++ b/mindspore/core/ops/lower_bound.cc @@ -15,6 +15,10 @@ */ #include "ops/lower_bound.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -60,6 +64,7 @@ TypePtr LowerBoundInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/lower_bound.h b/mindspore/core/ops/lower_bound.h index 9075b042fb..a33b815aea 100644 --- a/mindspore/core/ops/lower_bound.h +++ b/mindspore/core/ops/lower_bound.h @@ -22,23 +22,19 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "abstract/primitive_infer_map.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLowerBound = "LowerBound"; -class LowerBound : public PrimitiveC { +class MIND_API LowerBound : public BaseOperator { public: - LowerBound() : PrimitiveC(kNameLowerBound) { InitIOName({"sorted_x", "values"}, {"y"}); } - ~LowerBound() = default; - MS_DECLARE_PARENT(LowerBound, PrimitiveC); + MIND_API_BASE_MEMBER(LowerBound); + LowerBound() : BaseOperator(kNameLowerBound) { InitIOName({"sorted_x", "values"}, {"y"}); } }; -AbstractBasePtr LowerBoundInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LowerBoundInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLowerBound = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lp_norm.cc b/mindspore/core/ops/lp_norm.cc index f749a71ce8..6cb412ba09 100644 --- a/mindspore/core/ops/lp_norm.cc +++ b/mindspore/core/ops/lp_norm.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -98,6 +99,7 @@ TypePtr LpNormInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/lp_norm.h b/mindspore/core/ops/lp_norm.h index 2273dcac39..5f246ba8c8 100644 --- a/mindspore/core/ops/lp_norm.h +++ b/mindspore/core/ops/lp_norm.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLpNorm = "LpNorm"; -class LpNorm : public PrimitiveC { +class LpNorm : public BaseOperator { public: - LpNorm() : PrimitiveC(kNameLpNorm) { InitIOName({"input"}, {"output"}); } - ~LpNorm() = default; - MS_DECLARE_PARENT(LpNorm, PrimitiveC); + MIND_API_BASE_MEMBER(LpNorm); + LpNorm() : BaseOperator(kNameLpNorm) { InitIOName({"input"}, {"output"}); } }; -AbstractBasePtr LpNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LpNormInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLpNormPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/lp_normalization.cc b/mindspore/core/ops/lp_normalization.cc index c9214f01dd..83b9715c69 100644 --- a/mindspore/core/ops/lp_normalization.cc +++ b/mindspore/core/ops/lp_normalization.cc @@ -17,22 +17,24 @@ #include "ops/lp_normalization.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(LpNormalization, PrimitiveC, BaseOperator); void LpNormalization::Init(const int64_t axis, const int64_t p) { this->set_axis(axis); this->set_p(p); } -void LpNormalization::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void LpNormalization::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t LpNormalization::get_axis() const { auto value_ptr = this->GetAttr(kAxis); return GetValue(value_ptr); } -void LpNormalization::set_p(const int64_t p) { (void)this->AddAttr(kP, MakeValue(p)); } +void LpNormalization::set_p(const int64_t p) { (void)this->AddAttr(kP, api::MakeValue(p)); } int64_t LpNormalization::get_p() const { auto value_ptr = this->GetAttr(kP); diff --git a/mindspore/core/ops/lp_normalization.h b/mindspore/core/ops/lp_normalization.h index 9f84c7c28f..ed329e1b5a 100644 --- a/mindspore/core/ops/lp_normalization.h +++ b/mindspore/core/ops/lp_normalization.h @@ -18,23 +18,18 @@ #define MINDSPORE_CORE_OPS_LP_NORMALIZATION_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLpNormalization = "LpNormalization"; /// \brief LpNormalization defined LpNormalization operator prototype of lite. -class MS_CORE_API LpNormalization : public PrimitiveC { +class MIND_API LpNormalization : public BaseOperator { public: + MIND_API_BASE_MEMBER(LpNormalization); /// \brief Constructor. - LpNormalization() : PrimitiveC(kNameLpNormalization) {} - - /// \brief Destructor. - ~LpNormalization() = default; - - MS_DECLARE_PARENT(LpNormalization, PrimitiveC); + LpNormalization() : BaseOperator(kNameLpNormalization) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/lrn.cc b/mindspore/core/ops/lrn.cc index 2ae3d22fd9..72fe129d3e 100644 --- a/mindspore/core/ops/lrn.cc +++ b/mindspore/core/ops/lrn.cc @@ -23,12 +23,13 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void LRN::set_depth_radius(const int64_t depth_radius) { (void)CheckAndConvertUtils::CheckInteger(kDepthRadius, depth_radius, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kDepthRadius, MakeValue(depth_radius)); + (void)this->AddAttr(kDepthRadius, api::MakeValue(depth_radius)); } int64_t LRN::get_depth_radius() const { @@ -36,21 +37,21 @@ int64_t LRN::get_depth_radius() const { return GetValue(value_ptr); } -void LRN::set_bias(const float bias) { (void)this->AddAttr(kBias, MakeValue(bias)); } +void LRN::set_bias(const float bias) { (void)this->AddAttr(kBias, api::MakeValue(bias)); } float LRN::get_bias() const { auto value_ptr = GetAttr(kBias); return GetValue(value_ptr); } -void LRN::set_alpha(const float alpha) { (void)this->AddAttr(kAlpha, MakeValue(alpha)); } +void LRN::set_alpha(const float alpha) { (void)this->AddAttr(kAlpha, api::MakeValue(alpha)); } float LRN::get_alpha() const { auto value_ptr = GetAttr(kAlpha); return GetValue(value_ptr); } -void LRN::set_beta(const float beta) { (void)this->AddAttr(kBeta, MakeValue(beta)); } +void LRN::set_beta(const float beta) { (void)this->AddAttr(kBeta, api::MakeValue(beta)); } float LRN::get_beta() const { auto value_ptr = GetAttr(kBeta); @@ -58,7 +59,7 @@ float LRN::get_beta() const { } void LRN::set_norm_region(const std::string &norm_region) { CheckAndConvertUtils::CheckString(kNormRegion, norm_region, {"ACROSS_CHANNELS"}, this->name()); - (void)this->AddAttr(kNormRegion, MakeValue(norm_region)); + (void)this->AddAttr(kNormRegion, api::MakeValue(norm_region)); } std::string LRN::get_norm_region() const { @@ -74,6 +75,7 @@ void LRN::Init(const int64_t depth_radius, const float bias, const float alpha, this->set_norm_region(norm_region); } +MIND_API_BASE_IMPL(LRN, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameLRN, LRN); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lrn.h b/mindspore/core/ops/lrn.h index 1c8ff06adf..d670c9067c 100644 --- a/mindspore/core/ops/lrn.h +++ b/mindspore/core/ops/lrn.h @@ -20,21 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLRN = "LRN"; /// \brief Local Response Normalization. Refer to Python API @ref mindspore.ops.LRN for more details. -class MS_CORE_API LRN : public PrimitiveC { +class MIND_API LRN : public BaseOperator { public: + MIND_API_BASE_MEMBER(LRN); /// \brief Constructor. - LRN() : PrimitiveC(kNameLRN) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~LRN() = default; - MS_DECLARE_PARENT(LRN, PrimitiveC); + LRN() : BaseOperator(kNameLRN) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LRN for the inputs. void Init(const int64_t depth_radius = 5, const float bias = 1.0, const float alpha = 1.0, const float beta = 0.5, const std::string &norm_region = "ACROSS_CHANNELS"); @@ -69,8 +67,8 @@ class MS_CORE_API LRN : public PrimitiveC { /// \return norm_region. std::string get_norm_region() const; }; -AbstractBasePtr LrnInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LrnInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLrn = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lsh_projection.cc b/mindspore/core/ops/lsh_projection.cc index 7ea717a469..a849e98b23 100644 --- a/mindspore/core/ops/lsh_projection.cc +++ b/mindspore/core/ops/lsh_projection.cc @@ -16,14 +16,17 @@ #include "ops/lsh_projection.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(LshProjection, PrimitiveC, BaseOperator); void LshProjection::Init(const LshProjectionType &type) { set_type(type); } void LshProjection::set_type(const LshProjectionType &type) { int64_t swi = (int64_t)type; - (void)AddAttr(kType, MakeValue(swi)); + (void)AddAttr(kType, api::MakeValue(swi)); } LshProjectionType LshProjection::get_type() const { return LshProjectionType(GetValue(GetAttr(kType))); } diff --git a/mindspore/core/ops/lsh_projection.h b/mindspore/core/ops/lsh_projection.h index 6415f6239d..b399eda8d4 100644 --- a/mindspore/core/ops/lsh_projection.h +++ b/mindspore/core/ops/lsh_projection.h @@ -20,23 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLshProjection = "LshProjection"; /// \brief LshProjection defined LshProjection operator prototype of lite, which is to project an input to a bit vector. -class MS_CORE_API LshProjection : public PrimitiveC { +class MIND_API LshProjection : public BaseOperator { public: + MIND_API_BASE_MEMBER(LshProjection); /// \brief Constructor. - LshProjection() : PrimitiveC(kNameLshProjection) {} - - /// \brief Destructor. - ~LshProjection() = default; - - MS_DECLARE_PARENT(LshProjection, PrimitiveC); + LshProjection() : BaseOperator(kNameLshProjection) {} /// \brief Method to init the op's attributes. /// @@ -54,8 +50,8 @@ class MS_CORE_API LshProjection : public PrimitiveC { LshProjectionType get_type() const; }; -AbstractBasePtr LshProjectionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LshProjectionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lstm.cc b/mindspore/core/ops/lstm.cc index 71e6db4002..927eae2e46 100644 --- a/mindspore/core/ops/lstm.cc +++ b/mindspore/core/ops/lstm.cc @@ -15,6 +15,10 @@ */ #include "ops/lstm.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -72,48 +76,49 @@ AbstractBasePtr LstmInfer(const PrimitivePtr &primitive, const std::vectorname()); - (void)AddAttr(kInput_size, MakeValue(input_size)); + (void)AddAttr(kInput_size, api::MakeValue(input_size)); } int64_t LSTM::get_input_size() const { return GetValue(GetAttr(kInput_size)); } void LSTM::set_hidden_size(const int64_t hidden_size) { (void)CheckAndConvertUtils::CheckInteger(kHidden_size, hidden_size, kGreaterThan, 0, this->name()); - (void)AddAttr(kHidden_size, MakeValue(hidden_size)); + (void)AddAttr(kHidden_size, api::MakeValue(hidden_size)); } int64_t LSTM::get_hidden_size() const { return GetValue(GetAttr(kHidden_size)); } void LSTM::set_num_layers(const int64_t num_layers) { (void)CheckAndConvertUtils::CheckInteger(kNumLayers, num_layers, kGreaterThan, 0, this->name()); - (void)AddAttr(kNumLayers, MakeValue(num_layers)); + (void)AddAttr(kNumLayers, api::MakeValue(num_layers)); } int64_t LSTM::get_num_layers() const { return GetValue(GetAttr(kNumLayers)); } -void LSTM::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, MakeValue(has_bias)); } +void LSTM::set_has_bias(const bool has_bias) { (void)AddAttr(kHasBias, api::MakeValue(has_bias)); } bool LSTM::get_has_bias() const { auto value_ptr = this->GetAttr(kHasBias); return GetValue(value_ptr); } void LSTM::set_dropout(const float dropout) { CheckAndConvertUtils::CheckInRange(kDropout, dropout, kIncludeBoth, {0.0, 1.0}, this->name()); - (void)AddAttr(kDropout, MakeValue(dropout)); + (void)AddAttr(kDropout, api::MakeValue(dropout)); } float LSTM::get_dropout() const { auto value_ptr = this->GetAttr(kDropout); return GetValue(value_ptr); } -void LSTM::set_bidirectional(const bool bidirectional) { (void)AddAttr(kBidirectional, MakeValue(bidirectional)); } +void LSTM::set_bidirectional(const bool bidirectional) { (void)AddAttr(kBidirectional, api::MakeValue(bidirectional)); } bool LSTM::get_bidirectional() const { auto value_ptr = this->GetAttr(kBidirectional); return GetValue(value_ptr); } void LSTM::set_num_directions(const int64_t num_directions) { - (void)AddAttr(kNumDirections, MakeValue(num_directions)); + (void)AddAttr(kNumDirections, api::MakeValue(num_directions)); } int64_t LSTM::get_num_directions() const { return GetValue(GetAttr(kNumDirections)); } -void LSTM::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, MakeValue(zoneout_cell)); } +void LSTM::set_zoneout_cell(float zoneout_cell) { (void)AddAttr(kZoneoutCell, api::MakeValue(zoneout_cell)); } float LSTM::get_zoneout_cell() const { return GetValue(this->GetAttr(kZoneoutCell)); } -void LSTM::set_zoneout_hidden(float zoneout_hidden) { (void)AddAttr(kZoneoutHidden, MakeValue(zoneout_hidden)); } +void LSTM::set_zoneout_hidden(float zoneout_hidden) { (void)AddAttr(kZoneoutHidden, api::MakeValue(zoneout_hidden)); } float LSTM::get_zoneout_hidden() const { return GetValue(this->GetAttr(kZoneoutHidden)); } diff --git a/mindspore/core/ops/lstm.h b/mindspore/core/ops/lstm.h index c52a85d64e..bb8eb24a44 100644 --- a/mindspore/core/ops/lstm.h +++ b/mindspore/core/ops/lstm.h @@ -22,24 +22,20 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLSTM = "LSTM"; /// \brief Performs the Long Short-Term Memory (LSTM) on the input. /// Refer to Python API @ref mindspore.ops.LSTM for more details. -class MS_CORE_API LSTM : public PrimitiveC { +class MIND_API LSTM : public BaseOperator { public: + MIND_API_BASE_MEMBER(LSTM); /// \brief Constructor. - LSTM() : PrimitiveC(kNameLSTM) {} - /// \brief Destructor. - ~LSTM() = default; - MS_DECLARE_PARENT(LSTM, PrimitiveC); + LSTM() : BaseOperator(kNameLSTM) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.LSTM for the inputs. void Init(const int64_t input_size, const int64_t hidden_size, const int64_t num_layers, const bool has_bias, const float dropout, const bool bidirectional = false, const float zoneout_cell = 0.0f, @@ -103,8 +99,8 @@ class MS_CORE_API LSTM : public PrimitiveC { /// \return good_ld. int64_t get_good_ld(const int64_t dim, const int64_t type_size); }; -AbstractBasePtr LstmInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LstmInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lstsq.cc b/mindspore/core/ops/lstsq.cc index 9c9ec2700f..eec67d5de3 100644 --- a/mindspore/core/ops/lstsq.cc +++ b/mindspore/core/ops/lstsq.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -69,6 +70,7 @@ TypePtr LstsqInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/lstsq.h b/mindspore/core/ops/lstsq.h index 131c16caff..880149612f 100644 --- a/mindspore/core/ops/lstsq.h +++ b/mindspore/core/ops/lstsq.h @@ -20,21 +20,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLstsq = "Lstsq"; -class Lstsq : public PrimitiveC { +class MIND_API Lstsq : public BaseOperator { public: - Lstsq() : PrimitiveC(kNameLstsq) { InitIOName({"matrix", "rhs"}, {"y"}); } - ~Lstsq() = default; - MS_DECLARE_PARENT(Lstsq, PrimitiveC); + MIND_API_BASE_MEMBER(Lstsq); + Lstsq() : BaseOperator(kNameLstsq) { InitIOName({"matrix", "rhs"}, {"y"}); } }; -AbstractBasePtr LstsqInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LstsqInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLstsqPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/lu_solve_.cc b/mindspore/core/ops/lu_solve_.cc index af5feb76a7..cbda127712 100644 --- a/mindspore/core/ops/lu_solve_.cc +++ b/mindspore/core/ops/lu_solve_.cc @@ -17,6 +17,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" #define LuSolve_for(shape) \ do { \ @@ -147,6 +148,7 @@ TypePtr LuSolveInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/lu_solve_.h b/mindspore/core/ops/lu_solve_.h index 9d533cde8a..11289f88d3 100644 --- a/mindspore/core/ops/lu_solve_.h +++ b/mindspore/core/ops/lu_solve_.h @@ -22,21 +22,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameLuSolve = "LuSolve"; -class LuSolve : public PrimitiveC { +class MIND_API LuSolve : public BaseOperator { public: - LuSolve() : PrimitiveC(kNameLuSolve) { InitIOName({"x", "lu_data", "lu_pivots"}, {"output"}); } - ~LuSolve() = default; - MS_DECLARE_PARENT(LuSolve, PrimitiveC); + MIND_API_BASE_MEMBER(LuSolve); + LuSolve() : BaseOperator(kNameLuSolve) { InitIOName({"x", "lu_data", "lu_pivots"}, {"output"}); } }; -AbstractBasePtr LuSolveInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr LuSolveInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimLuSolvePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/make_tuple.cc b/mindspore/core/ops/make_tuple.cc index e82324cc09..8e05dd5ece 100644 --- a/mindspore/core/ops/make_tuple.cc +++ b/mindspore/core/ops/make_tuple.cc @@ -15,9 +15,12 @@ */ #include "ops/make_tuple.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(MakeTuple, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMakeTuple, MakeTuple); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/make_tuple.h b/mindspore/core/ops/make_tuple.h index 0b560cc574..ea4bf5a670 100644 --- a/mindspore/core/ops/make_tuple.h +++ b/mindspore/core/ops/make_tuple.h @@ -16,21 +16,17 @@ #ifndef MINDSPORE_CORE_OPS_MAKE_TUPLE_H_ #define MINDSPORE_CORE_OPS_MAKE_TUPLE_H_ -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameMakeTuple = "MakeTuple"; /// \brief MakeTuple op is used to pack multiple nodes into a whole, which is only used in FuncGraph. -class MS_CORE_API MakeTuple : public PrimitiveC { +class MIND_API MakeTuple : public BaseOperator { public: + MIND_API_BASE_MEMBER(MakeTuple); /// \brief Constructor. - MakeTuple() : PrimitiveC(kNameMakeTuple) {} - - /// \brief Destructor. - ~MakeTuple() = default; - - MS_DECLARE_PARENT(MakeTuple, PrimitiveC); + MakeTuple() : BaseOperator(kNameMakeTuple) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/masked_fill.cc b/mindspore/core/ops/masked_fill.cc index f7d207703d..d6289aa050 100644 --- a/mindspore/core/ops/masked_fill.cc +++ b/mindspore/core/ops/masked_fill.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -69,6 +70,7 @@ TypePtr MaskedFillInferType(const PrimitivePtr &prim, const std::vector &input_args) { return std::make_shared(MaskedFillInferType(primitive, input_args), diff --git a/mindspore/core/ops/masked_fill.h b/mindspore/core/ops/masked_fill.h index de12987a0a..08d56f02a2 100644 --- a/mindspore/core/ops/masked_fill.h +++ b/mindspore/core/ops/masked_fill.h @@ -19,27 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMaskedFill = "MaskedFill"; /// \brief Fills elements of self tensor with value where mask is True. /// Refer to Python API @ref mindspore.ops.MaskedFill for more details. -class MaskedFill : public PrimitiveC { +class MaskedFill : public BaseOperator { public: + MIND_API_BASE_MEMBER(MaskedFill); /// \brief Constructor. - MaskedFill() : PrimitiveC(kNameMaskedFill) { InitIOName({"input", "mask", "value"}, {"output"}); } - /// \brief Destructor. - ~MaskedFill() = default; - MS_DECLARE_PARENT(MaskedFill, PrimitiveC); + MaskedFill() : BaseOperator(kNameMaskedFill) { InitIOName({"input", "mask", "value"}, {"output"}); } }; -AbstractBasePtr MaskedFillInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MaskedFillInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_MASKED_FILL_H_ diff --git a/mindspore/core/ops/mat_mul.cc b/mindspore/core/ops/mat_mul.cc index aaef5abf8e..0eb79cf267 100644 --- a/mindspore/core/ops/mat_mul.cc +++ b/mindspore/core/ops/mat_mul.cc @@ -18,6 +18,9 @@ #include #include #include "ops/mat_mul.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -26,9 +29,9 @@ void MatMul::Init(bool transpose_a, bool transpose_b) { set_transpose_b(transpose_b); } -void MatMul::set_transpose_a(bool transpose_a) { (void)AddAttr(kTransposeA, MakeValue(transpose_a)); } +void MatMul::set_transpose_a(bool transpose_a) { (void)AddAttr(kTransposeA, api::MakeValue(transpose_a)); } -void MatMul::set_transpose_b(bool transpose_b) { (void)AddAttr(kTransposeB, MakeValue(transpose_b)); } +void MatMul::set_transpose_b(bool transpose_b) { (void)AddAttr(kTransposeB, api::MakeValue(transpose_b)); } bool MatMul::get_transpose_a() const { auto value_ptr = GetAttr(kTransposeA); @@ -40,6 +43,7 @@ bool MatMul::get_transpose_b() const { return GetValue(value_ptr); } +MIND_API_BASE_IMPL(MatMul, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMatMul, MatMul); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mat_mul.h b/mindspore/core/ops/mat_mul.h index 054899319c..6d076bc94c 100644 --- a/mindspore/core/ops/mat_mul.h +++ b/mindspore/core/ops/mat_mul.h @@ -19,24 +19,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "abstract/dshape.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMatMul = "MatMul"; /// \brief Multiplies matrix a and matrix b. Refer to Python API @ref mindspore.ops.MatMul for more details. -class MS_CORE_API MatMul : public PrimitiveC { +class MIND_API MatMul : public BaseOperator { public: + MIND_API_BASE_MEMBER(MatMul); /// \brief Constructor. - MatMul() : PrimitiveC(kNameMatMul) { InitIOName({"x1", "x2"}, {"output"}); } - explicit MatMul(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "x2"}, {"output"}); } - /// \brief Destructor. - ~MatMul() = default; - MS_DECLARE_PARENT(MatMul, PrimitiveC); + MatMul() : BaseOperator(kNameMatMul) { InitIOName({"x1", "x2"}, {"output"}); } + explicit MatMul(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "x2"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.MatMul for the inputs. void Init(bool transpose_a = false, bool transpose_b = false); /// \brief Set transpose_a. diff --git a/mindspore/core/ops/matrix_determinant.cc b/mindspore/core/ops/matrix_determinant.cc index 78463addb8..326d82bf20 100644 --- a/mindspore/core/ops/matrix_determinant.cc +++ b/mindspore/core/ops/matrix_determinant.cc @@ -20,6 +20,8 @@ #include "ops/op_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -47,6 +49,7 @@ TypePtr MatrixDeterminantInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/matrix_determinant.h b/mindspore/core/ops/matrix_determinant.h index bebed78869..46e730458e 100644 --- a/mindspore/core/ops/matrix_determinant.h +++ b/mindspore/core/ops/matrix_determinant.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMatrixDeterminant = "MatrixDeterminant"; -class MatrixDeterminant : public PrimitiveC { +class MIND_API MatrixDeterminant : public BaseOperator { public: - MatrixDeterminant() : PrimitiveC(kNameMatrixDeterminant) { InitIOName({"x"}, {"y"}); } - ~MatrixDeterminant() = default; - MS_DECLARE_PARENT(MatrixDeterminant, PrimitiveC); + MIND_API_BASE_MEMBER(MatrixDeterminant); + MatrixDeterminant() : BaseOperator(kNameMatrixDeterminant) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr MatrixDeterminantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MatrixDeterminantInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimMatrixDeterminantPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/matrix_diag_part.cc b/mindspore/core/ops/matrix_diag_part.cc index a06fd0bb6d..7d1ec7b8ed 100644 --- a/mindspore/core/ops/matrix_diag_part.cc +++ b/mindspore/core/ops/matrix_diag_part.cc @@ -19,6 +19,8 @@ #include "abstract/primitive_infer_map.h" #include "utils/check_convert_utils.h" #include "abstract/utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -50,6 +52,7 @@ TypePtr MatrixDiagPartInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(MatrixDiagPartInferShape(primitive, input_args), diff --git a/mindspore/core/ops/matrix_diag_part.h b/mindspore/core/ops/matrix_diag_part.h index 75bcaceb51..4ba8f70912 100644 --- a/mindspore/core/ops/matrix_diag_part.h +++ b/mindspore/core/ops/matrix_diag_part.h @@ -19,27 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMatrixDiagPart = "MatrixDiagPartV3"; /// \brief get the specified part of the inner most diag matrix of a matrix, fill with padding value . /// Refer to Python API @ref mindspore.ops.MatrixDiagPart for more details. -class MatrixDiagPartV3 : public PrimitiveC { +class MatrixDiagPartV3 : public BaseOperator { public: + MIND_API_BASE_MEMBER(MatrixDiagPartV3); /// \brief Constructor. - MatrixDiagPartV3() : PrimitiveC(kNameMatrixDiagPart) { InitIOName({"input", "k", "padding_value"}, {"output"}); } - /// \brief Destructor. - ~MatrixDiagPartV3() = default; - MS_DECLARE_PARENT(MatrixDiagPartV3, PrimitiveC); + MatrixDiagPartV3() : BaseOperator(kNameMatrixDiagPart) { InitIOName({"input", "k", "padding_value"}, {"output"}); } }; -AbstractBasePtr MatrixDiagPartInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MatrixDiagPartInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_MATRIX_DIAG_PART_H_ diff --git a/mindspore/core/ops/matrix_inverse.cc b/mindspore/core/ops/matrix_inverse.cc index b5c039ccd2..972b83bf88 100644 --- a/mindspore/core/ops/matrix_inverse.cc +++ b/mindspore/core/ops/matrix_inverse.cc @@ -20,6 +20,8 @@ #include "ops/op_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -46,6 +48,7 @@ TypePtr MatrixInverseInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/matrix_inverse.h b/mindspore/core/ops/matrix_inverse.h index ac0b936b90..0c2071387d 100644 --- a/mindspore/core/ops/matrix_inverse.h +++ b/mindspore/core/ops/matrix_inverse.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMatrixInverse = "MatrixInverse"; -class MatrixInverse : public PrimitiveC { +class MIND_API MatrixInverse : public BaseOperator { public: - MatrixInverse() : PrimitiveC(kNameMatrixInverse) { InitIOName({"x"}, {"y"}); } - ~MatrixInverse() = default; - MS_DECLARE_PARENT(MatrixInverse, PrimitiveC); + MIND_API_BASE_MEMBER(MatrixInverse); + MatrixInverse() : BaseOperator(kNameMatrixInverse) { InitIOName({"x"}, {"y"}); } }; -AbstractBasePtr MatrixInverseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MatrixInverseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimMatrixInversePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/max_pool.cc b/mindspore/core/ops/max_pool.cc index 80e253aabe..58dec58b57 100644 --- a/mindspore/core/ops/max_pool.cc +++ b/mindspore/core/ops/max_pool.cc @@ -23,35 +23,37 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void MaxPool::set_pad_mode(const PadMode &pad_mode) { int64_t swi = pad_mode; - (void)this->AddAttr(kPadMode, MakeValue(swi)); + (void)this->AddAttr(kPadMode, api::MakeValue(swi)); } PadMode MaxPool::get_pad_mode() const { return PadMode(GetValue(GetAttr(kPadMode))); } void MaxPool::set_kernel_size(const std::vector &kernel_size) { - (void)this->AddAttr(kKernelSize, - MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, this->name()))); + (void)this->AddAttr( + kKernelSize, api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kKernelSize, kernel_size, this->name()))); } std::vector MaxPool::get_kernel_size() const { return GetValue>(GetAttr(kKernelSize)); } void MaxPool::set_strides(const std::vector &strides) { - (void)this->AddAttr(kStrides, MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStrides, strides, this->name()))); + (void)this->AddAttr(kStrides, + api::MakeValue(CheckAndConvertUtils::CheckPositiveVector(kStrides, strides, this->name()))); } std::vector MaxPool::get_strides() const { return GetValue>(GetAttr(kStrides)); } void MaxPool::set_format(const Format &format) { int64_t f = format; - (void)this->AddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } Format MaxPool::get_format() const { return Format(GetValue(GetAttr(kFormat))); } -void MaxPool::set_pad(const std::vector &pad) { (void)this->AddAttr(kPad, MakeValue(pad)); } +void MaxPool::set_pad(const std::vector &pad) { (void)this->AddAttr(kPad, api::MakeValue(pad)); } std::vector MaxPool::get_pad() const { auto value_ptr = GetAttr(kPad); @@ -60,7 +62,7 @@ std::vector MaxPool::get_pad() const { void MaxPool::set_round_mode(const RoundMode &round_mode) { int64_t swi = round_mode; - (void)this->AddAttr(kRoundMode, MakeValue(swi)); + (void)this->AddAttr(kRoundMode, api::MakeValue(swi)); } RoundMode MaxPool::get_round_mode() const { @@ -78,6 +80,7 @@ void MaxPool::Init(const std::vector &kernel_size, const std::vectorset_round_mode(round_mode); } +MIND_API_BASE_IMPL(MaxPool, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMaxPool, MaxPool); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/max_pool.h b/mindspore/core/ops/max_pool.h index 7d3352ada4..c750028353 100644 --- a/mindspore/core/ops/max_pool.h +++ b/mindspore/core/ops/max_pool.h @@ -21,22 +21,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameMaxPool = "MaxPool"; /// \brief Max pooling operation. Refer to Python API @ref mindspore.ops.MaxPool for more details. -class MS_CORE_API MaxPool : public PrimitiveC { +class MIND_API MaxPool : public BaseOperator { public: + MIND_API_BASE_MEMBER(MaxPool); /// \brief Constructor. - MaxPool() : PrimitiveC(kNameMaxPool) { InitIOName({"x"}, {"output"}); } - explicit MaxPool(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~MaxPool() = default; - MS_DECLARE_PARENT(MaxPool, PrimitiveC); + MaxPool() : BaseOperator(kNameMaxPool) { InitIOName({"x"}, {"output"}); } + explicit MaxPool(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.MaxPool for the inputs. void Init(const std::vector &kernel_size = {1}, const std::vector &stride = {1}, const PadMode &pad_mode = VALID, const Format &format = NCHW, @@ -80,8 +79,8 @@ class MS_CORE_API MaxPool : public PrimitiveC { RoundMode get_round_mode() const; }; -AbstractBasePtr MaxPoolInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MaxPoolInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/maximum.cc b/mindspore/core/ops/maximum.cc index 8a8a6ac9fe..cd8956d233 100644 --- a/mindspore/core/ops/maximum.cc +++ b/mindspore/core/ops/maximum.cc @@ -18,6 +18,8 @@ #include #include "ops/maximum.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -75,6 +77,8 @@ TypePtr MaximumInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = MaximumInferType(primitive, input_args); diff --git a/mindspore/core/ops/maximum.h b/mindspore/core/ops/maximum.h index 47ac12836c..77ef8fa6be 100644 --- a/mindspore/core/ops/maximum.h +++ b/mindspore/core/ops/maximum.h @@ -20,25 +20,23 @@ #include #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameMaximum = "Maximum"; /// \brief Computes the maximum of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.Maximum for more details. -class MS_CORE_API Maximum : public PrimitiveC { +class MIND_API Maximum : public BaseOperator { public: + MIND_API_BASE_MEMBER(Maximum); /// \brief Constructor. - Maximum() : PrimitiveC(kNameMaximum) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Maximum() = default; - MS_DECLARE_PARENT(Maximum, PrimitiveC); + Maximum() : BaseOperator(kNameMaximum) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Maximum for the inputs. void Init() const {} }; -AbstractBasePtr MaximumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MaximumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimMaximumPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mfcc.cc b/mindspore/core/ops/mfcc.cc index 5580c522fa..b60386b85f 100644 --- a/mindspore/core/ops/mfcc.cc +++ b/mindspore/core/ops/mfcc.cc @@ -18,6 +18,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -30,7 +31,7 @@ void Mfcc::Init(const float freq_upper_limit, const float freq_lower_limit, cons } void Mfcc::set_freq_upper_limit(const float freq_upper_limit) { - (void)this->AddAttr(kFreqUpperLimit, MakeValue(freq_upper_limit)); + (void)this->AddAttr(kFreqUpperLimit, api::MakeValue(freq_upper_limit)); } float Mfcc::get_freq_upper_limit() const { @@ -39,7 +40,7 @@ float Mfcc::get_freq_upper_limit() const { } void Mfcc::set_freq_lower_limit(const float freq_lower_limit) { - (void)this->AddAttr(kFreqLowerLimit, MakeValue(freq_lower_limit)); + (void)this->AddAttr(kFreqLowerLimit, api::MakeValue(freq_lower_limit)); } float Mfcc::get_freq_lower_limit() const { @@ -48,7 +49,7 @@ float Mfcc::get_freq_lower_limit() const { } void Mfcc::set_filter_bank_channel_num(const int64_t filter_bank_channel_num) { - (void)this->AddAttr(kFilterBankChannelNum, MakeValue(filter_bank_channel_num)); + (void)this->AddAttr(kFilterBankChannelNum, api::MakeValue(filter_bank_channel_num)); } int64_t Mfcc::get_filter_bank_channel_num() const { @@ -57,11 +58,12 @@ int64_t Mfcc::get_filter_bank_channel_num() const { } void Mfcc::set_dct_coeff_num(const int64_t dct_coeff_num) { - (void)this->AddAttr(kDctCoeffNum, MakeValue(dct_coeff_num)); + (void)this->AddAttr(kDctCoeffNum, api::MakeValue(dct_coeff_num)); } int64_t Mfcc::get_dct_coeff_num() const { return GetValue(GetAttr(kDctCoeffNum)); } +MIND_API_BASE_IMPL(Mfcc, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMfcc, Mfcc); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mfcc.h b/mindspore/core/ops/mfcc.h index d4f7c757b5..1b959e564a 100644 --- a/mindspore/core/ops/mfcc.h +++ b/mindspore/core/ops/mfcc.h @@ -18,23 +18,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMfcc = "Mfcc"; /// \brief Mfcc defined the operator prototype of extracting Mel-Frequency Cepstral Coefficients. -class MS_CORE_API Mfcc : public PrimitiveC { +class MIND_API Mfcc : public BaseOperator { public: + MIND_API_BASE_MEMBER(Mfcc); /// \brief Constructor. - Mfcc() : PrimitiveC(kNameMfcc) {} - - /// \brief Destructor. - ~Mfcc() = default; - - MS_DECLARE_PARENT(Mfcc, PrimitiveC); + Mfcc() : BaseOperator(kNameMfcc) {} /// \brief Method to init the op's attributes. /// @@ -85,8 +80,8 @@ class MS_CORE_API Mfcc : public PrimitiveC { /// \return the output channels to generate per time slice. int64_t get_dct_coeff_num() const; }; -AbstractBasePtr MfccInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MfccInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/minimum.cc b/mindspore/core/ops/minimum.cc index 5bb91e7f5b..2f0bd70881 100644 --- a/mindspore/core/ops/minimum.cc +++ b/mindspore/core/ops/minimum.cc @@ -23,9 +23,11 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Minimum, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMinimum, Minimum); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/minimum.h b/mindspore/core/ops/minimum.h index 5d1e4a7b1c..22b6e789d6 100644 --- a/mindspore/core/ops/minimum.h +++ b/mindspore/core/ops/minimum.h @@ -20,28 +20,26 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMinimum = "Minimum"; /// \brief Computes the minimum of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.Minimum for more details. -class MS_CORE_API Minimum : public PrimitiveC { +class MIND_API Minimum : public BaseOperator { public: + MIND_API_BASE_MEMBER(Minimum); /// \brief Constructor. - Minimum() : PrimitiveC(kNameMinimum) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Minimum() = default; - MS_DECLARE_PARENT(Minimum, PrimitiveC); + Minimum() : BaseOperator(kNameMinimum) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Minimum for the inputs. void Init() const {} }; -AbstractBasePtr MinimumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MinimumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mod.cc b/mindspore/core/ops/mod.cc index dd2b9ee434..408a4e41a6 100644 --- a/mindspore/core/ops/mod.cc +++ b/mindspore/core/ops/mod.cc @@ -25,6 +25,8 @@ #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -82,6 +84,7 @@ AbstractBasePtr ModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr auto infer_shape = ModInferShape(primitive, input_args); return abstract::MakeAbstract(infer_shape, infer_type); } +MIND_API_BASE_IMPL(Mod, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameMod, Mod); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mod.h b/mindspore/core/ops/mod.h index 77bbef9e67..43ab1bf848 100644 --- a/mindspore/core/ops/mod.h +++ b/mindspore/core/ops/mod.h @@ -18,27 +18,24 @@ #define MINDSPORE_CORE_OPS_MOD_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameMod = "Mod"; /// \brief Computes the remainder of dividing the first input tensor by the second input tensor element-wise. /// Refer to Python API @ref mindspore.ops.Mod for more details. -class MS_CORE_API Mod : public PrimitiveC { +class MIND_API Mod : public BaseOperator { public: + MIND_API_BASE_MEMBER(Mod); /// \brief Constructor. - Mod() : PrimitiveC(kNameMod) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Mod() = default; - MS_DECLARE_PARENT(Mod, PrimitiveC); + Mod() : BaseOperator(kNameMod) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Mod for the inputs. void Init() const {} }; -AbstractBasePtr ModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimModPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mul.cc b/mindspore/core/ops/mul.cc index 8763cba033..98ff1c8ba6 100644 --- a/mindspore/core/ops/mul.cc +++ b/mindspore/core/ops/mul.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -52,6 +53,7 @@ TypePtr MulInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(MulInferShape(primitive, input_args), MulInferType(primitive, input_args)); diff --git a/mindspore/core/ops/mul.h b/mindspore/core/ops/mul.h index 8293935d40..eee019d7af 100644 --- a/mindspore/core/ops/mul.h +++ b/mindspore/core/ops/mul.h @@ -20,27 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameMul = prim::kMul; +constexpr auto kNameMul = "Mul"; /// \brief Multiply two tensors element-wise. Refer to Python API @ref mindspore.ops.Mul for more details. -class MS_CORE_API Mul : public PrimitiveC { +class MIND_API Mul : public BaseOperator { public: + MIND_API_BASE_MEMBER(Mul); /// \brief Constructor. - Mul() : PrimitiveC(kNameMul) { InitIOName({"x", "y"}, {"output"}); } - explicit Mul(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Mul() = default; - MS_DECLARE_PARENT(Mul, PrimitiveC); + Mul() : BaseOperator(kNameMul) { InitIOName({"x", "y"}, {"output"}); } + explicit Mul(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Mul for the inputs. void Init() const {} }; -AbstractBasePtr MulInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MulInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/mulnonan.cc b/mindspore/core/ops/mulnonan.cc index 30f9d0367d..b1e84f2941 100644 --- a/mindspore/core/ops/mulnonan.cc +++ b/mindspore/core/ops/mulnonan.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -74,6 +75,7 @@ TypePtr MulNoNanInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infer_type = MulNoNanInferType(primitive, input_args); diff --git a/mindspore/core/ops/mulnonan.h b/mindspore/core/ops/mulnonan.h index f121bb0905..8b21622c77 100644 --- a/mindspore/core/ops/mulnonan.h +++ b/mindspore/core/ops/mulnonan.h @@ -20,23 +20,21 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameMulNoNan = prim::kMulNoNan; -class MulNoNan : public PrimitiveC { +constexpr auto kNameMulNoNan = "MulNoNan"; +class MulNoNan : public BaseOperator { public: - MulNoNan() : PrimitiveC(kNameMulNoNan) { InitIOName({"x", "y"}, {"output"}); } - explicit MulNoNan(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "y"}, {"output"}); } - ~MulNoNan() = default; - MS_DECLARE_PARENT(MulNoNan, PrimitiveC); + MIND_API_BASE_MEMBER(MulNoNan); + MulNoNan() : BaseOperator(kNameMulNoNan) { InitIOName({"x", "y"}, {"output"}); } + explicit MulNoNan(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "y"}, {"output"}); } void Init() {} }; -AbstractBasePtr MulNoNanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr MulNoNanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimMulNoNanPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/neg.cc b/mindspore/core/ops/neg.cc index 4a4c81f9f3..eeb1baa457 100644 --- a/mindspore/core/ops/neg.cc +++ b/mindspore/core/ops/neg.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -138,6 +139,7 @@ ValuePtr NegInferValue(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/neg.h b/mindspore/core/ops/neg.h index 77d32ad686..c5fd5e14fc 100644 --- a/mindspore/core/ops/neg.h +++ b/mindspore/core/ops/neg.h @@ -18,28 +18,25 @@ #define MINDSPORE_CORE_OPS_NEG_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameNeg = prim::kNeg; +constexpr auto kNameNeg = "Neg"; /// \brief Returns a tensor with negative values of the input tensor element-wise. /// Refer to Python API @ref mindspore.ops.Neg for more details. -class MS_CORE_API Neg : public PrimitiveC { +class MIND_API Neg : public BaseOperator { public: + MIND_API_BASE_MEMBER(Neg); /// \brief Constructor. - Neg() : PrimitiveC(prim::kPrimNeg->name()) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Neg() = default; - MS_DECLARE_PARENT(Neg, PrimitiveC); + Neg() : BaseOperator("Neg") { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Neg for the inputs. void Init() const {} }; -AbstractBasePtr NegInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr NegInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/neighborexchange.cc b/mindspore/core/ops/neighborexchange.cc index 2a4083a08d..4ddac712f4 100644 --- a/mindspore/core/ops/neighborexchange.cc +++ b/mindspore/core/ops/neighborexchange.cc @@ -18,6 +18,7 @@ #include #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -187,6 +188,8 @@ TypePtr NeighborExchangeInferType(const PrimitivePtr &primitive) { return std::make_shared(type_vec); } } // namespace + +MIND_API_BASE_IMPL(NeighborExchange, PrimitiveC, BaseOperator); AbstractBasePtr NeighborExchangeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { NeighborExchangeCheck(primitive, input_args); diff --git a/mindspore/core/ops/neighborexchange.h b/mindspore/core/ops/neighborexchange.h index e70ce19642..165ac9a2dd 100644 --- a/mindspore/core/ops/neighborexchange.h +++ b/mindspore/core/ops/neighborexchange.h @@ -18,27 +18,24 @@ #define MINDSPORE_CORE_OPS_NEIGHBOREXCHANGE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNeighborExchange = "NeighborExchange"; /// \brief NeighborExchange sends data from the local rank to ranks in the send_rank_ids. /// Refer to Python API @ref mindspore.ops.NeighborExchange for more details. -class MS_CORE_API NeighborExchange : public PrimitiveC { +class MIND_API NeighborExchange : public BaseOperator { public: + MIND_API_BASE_MEMBER(NeighborExchange); /// \brief Constructor. - NeighborExchange() : PrimitiveC(kNameNeighborExchange) {} - /// \brief Destructor. - ~NeighborExchange() = default; - MS_DECLARE_PARENT(NeighborExchange, PrimitiveC); + NeighborExchange() : BaseOperator(kNameNeighborExchange) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.NeighborExchange for the inputs. void Init() const {} }; -AbstractBasePtr NeighborExchangeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr NeighborExchangeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/neighborexchangev2.cc b/mindspore/core/ops/neighborexchangev2.cc index 1a7b2c9d64..da85418d36 100644 --- a/mindspore/core/ops/neighborexchangev2.cc +++ b/mindspore/core/ops/neighborexchangev2.cc @@ -19,6 +19,7 @@ #include #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -230,6 +231,8 @@ TypePtr NeighborExchangeV2InferType(const PrimitivePtr &primitive, const std::ve return recv_type; } } // namespace + +MIND_API_BASE_IMPL(NeighborExchangeV2, PrimitiveC, BaseOperator); AbstractBasePtr NeighborExchangeV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { NeighborExchangeV2Check(primitive, input_args); diff --git a/mindspore/core/ops/neighborexchangev2.h b/mindspore/core/ops/neighborexchangev2.h index 0da412c1f9..c72e71ece2 100644 --- a/mindspore/core/ops/neighborexchangev2.h +++ b/mindspore/core/ops/neighborexchangev2.h @@ -18,24 +18,22 @@ #define MINDSPORE_CORE_OPS_NEIGHBOREXCHANGEV2_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNeighborExchangeV2 = "NeighborExchangeV2"; -class MS_CORE_API NeighborExchangeV2 : public PrimitiveC { +class MIND_API NeighborExchangeV2 : public BaseOperator { public: - NeighborExchangeV2() : PrimitiveC(kNameNeighborExchangeV2) {} - ~NeighborExchangeV2() = default; - MS_DECLARE_PARENT(NeighborExchangeV2, PrimitiveC); + MIND_API_BASE_MEMBER(NeighborExchangeV2); + NeighborExchangeV2() : BaseOperator(kNameNeighborExchangeV2) {} void Init() {} }; using kPrimNeighborExchangeV2Ptr = std::shared_ptr; -AbstractBasePtr NeighborExchangeV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr NeighborExchangeV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/nllloss.cc b/mindspore/core/ops/nllloss.cc index 9916d0f66a..3432ff804d 100644 --- a/mindspore/core/ops/nllloss.cc +++ b/mindspore/core/ops/nllloss.cc @@ -17,6 +17,7 @@ #include "ops/nllloss.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -24,7 +25,7 @@ void NLLLoss::Init(const Reduction &reduction) { set_reduction(reduction); } void NLLLoss::set_reduction(const Reduction &reduction) { int64_t reduce = reduction; - (void)AddAttr(kReduction, MakeValue(reduce)); + (void)AddAttr(kReduction, api::MakeValue(reduce)); } Reduction NLLLoss::get_reduction() const { @@ -32,6 +33,7 @@ Reduction NLLLoss::get_reduction() const { return Reduction(GetValue(value_ptr)); } +MIND_API_BASE_IMPL(NLLLoss, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameNLLLoss, NLLLoss) } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/nllloss.h b/mindspore/core/ops/nllloss.h index de2fb512e5..912a11ffd1 100644 --- a/mindspore/core/ops/nllloss.h +++ b/mindspore/core/ops/nllloss.h @@ -19,22 +19,18 @@ #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" #include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNLLLoss = "NLLLoss"; /// \brief NLLLoss operation. Refer to Python API @ref mindspore.ops.NLLLoss for more details. -class MS_CORE_API NLLLoss : public PrimitiveC { +class MIND_API NLLLoss : public BaseOperator { public: + MIND_API_BASE_MEMBER(NLLLoss); /// \brief Constructor. - NLLLoss() : PrimitiveC(kNameNLLLoss) { InitIOName({"logits", "labels", "weight"}, {"loss", "total_weight"}); } - - /// \brief Destructor. - ~NLLLoss() = default; - - MS_DECLARE_PARENT(NLLLoss, PrimitiveC); + NLLLoss() : BaseOperator(kNameNLLLoss) { InitIOName({"logits", "labels", "weight"}, {"loss", "total_weight"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.NLLLoss for the inputs. void Init(const Reduction &reduction = NONE); diff --git a/mindspore/core/ops/non_max_suppression.cc b/mindspore/core/ops/non_max_suppression.cc index 132cbbc71a..296433448f 100644 --- a/mindspore/core/ops/non_max_suppression.cc +++ b/mindspore/core/ops/non_max_suppression.cc @@ -17,11 +17,16 @@ #include #include "ops/non_max_suppression.h" +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(NonMaxSuppression, PrimitiveC, BaseOperator); void NonMaxSuppression::set_center_point_box(const int64_t center_point_box) { - (void)AddAttr(kCenterPointBox, MakeValue(center_point_box)); + (void)AddAttr(kCenterPointBox, api::MakeValue(center_point_box)); } int64_t NonMaxSuppression::get_center_point_box() const { auto value_ptr = this->GetAttr(kCenterPointBox); diff --git a/mindspore/core/ops/non_max_suppression.h b/mindspore/core/ops/non_max_suppression.h index 0c3c8682ed..8f2fe47f0c 100644 --- a/mindspore/core/ops/non_max_suppression.h +++ b/mindspore/core/ops/non_max_suppression.h @@ -22,25 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNonMaxSuppression = "NonMaxSuppression"; /// \brief NonMaxSuppression QuantDTypeCast the NonMaxSuppression operator prototype. -class MS_CORE_API NonMaxSuppression : public PrimitiveC { +class MIND_API NonMaxSuppression : public BaseOperator { public: + MIND_API_BASE_MEMBER(NonMaxSuppression); /// \brief Constructor. - NonMaxSuppression() : PrimitiveC(kNameNonMaxSuppression) {} - - /// \brief Destructor. - ~NonMaxSuppression() = default; - - MS_DECLARE_PARENT(NonMaxSuppression, PrimitiveC); + NonMaxSuppression() : BaseOperator(kNameNonMaxSuppression) {} /// \brief Method to init the op's attributes. /// @@ -59,8 +53,8 @@ class MS_CORE_API NonMaxSuppression : public PrimitiveC { /// \return an integer value. int64_t get_center_point_box() const; }; -AbstractBasePtr NonMaxSuppressionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr NonMaxSuppressionInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimNonMaxSuppressionPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/non_max_suppression_v3.cc b/mindspore/core/ops/non_max_suppression_v3.cc index 6c31a07e17..e357cb3d1f 100644 --- a/mindspore/core/ops/non_max_suppression_v3.cc +++ b/mindspore/core/ops/non_max_suppression_v3.cc @@ -17,6 +17,10 @@ #include #include "ops/non_max_suppression_v3.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "abstract/dshape.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -105,6 +109,8 @@ TypePtr NonMaxSuppressionV3InferType(const PrimitivePtr &prim, const std::vector return max_output_size_type; } } // namespace + +MIND_API_BASE_IMPL(NonMaxSuppressionV3, PrimitiveC, BaseOperator); AbstractBasePtr NonMaxSuppressionV3Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/non_max_suppression_v3.h b/mindspore/core/ops/non_max_suppression_v3.h index d7347745f3..325136974c 100644 --- a/mindspore/core/ops/non_max_suppression_v3.h +++ b/mindspore/core/ops/non_max_suppression_v3.h @@ -22,26 +22,22 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "abstract/dshape.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNonMaxSuppressionV3 = "NonMaxSuppressionV3"; -class NonMaxSuppressionV3 : public PrimitiveC { +class NonMaxSuppressionV3 : public BaseOperator { public: - NonMaxSuppressionV3() : PrimitiveC(kNameNonMaxSuppressionV3) { + MIND_API_BASE_MEMBER(NonMaxSuppressionV3); + NonMaxSuppressionV3() : BaseOperator(kNameNonMaxSuppressionV3) { InitIOName({"boxes", "score", "max_output_size", "iou_threshold", "score_threshold"}, {"selected_indices"}); } - ~NonMaxSuppressionV3() = default; - MS_DECLARE_PARENT(NonMaxSuppressionV3, PrimitiveC); }; -AbstractBasePtr NonMaxSuppressionV3Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr NonMaxSuppressionV3Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimNonMaxSuppressionV3Ptr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/non_zero.cc b/mindspore/core/ops/non_zero.cc index aedc8ef245..4d10b2ac97 100644 --- a/mindspore/core/ops/non_zero.cc +++ b/mindspore/core/ops/non_zero.cc @@ -16,9 +16,12 @@ #include #include "ops/non_zero.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(NonZero, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameNonZero, NonZero); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/non_zero.h b/mindspore/core/ops/non_zero.h index 24891571ed..1790e388cb 100644 --- a/mindspore/core/ops/non_zero.h +++ b/mindspore/core/ops/non_zero.h @@ -17,22 +17,19 @@ #ifndef MINDSPORE_CORE_OPS_NON_ZERO_H_ #define MINDSPORE_CORE_OPS_NON_ZERO_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameNonZero = "NonZero"; /// \brief Calculate tensor not zero index, by default. /// Refer to Python API @ref mindspore.ops.NonZero for more details. -class MS_CORE_API NonZero : public PrimitiveC { +class MIND_API NonZero : public BaseOperator { public: + MIND_API_BASE_MEMBER(NonZero); /// \brief Constructor. - NonZero() : PrimitiveC(kNameNonZero) {} - /// \brief Destructor. - ~NonZero() = default; - MS_DECLARE_PARENT(NonZero, PrimitiveC); + NonZero() : BaseOperator(kNameNonZero) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/not_equal.cc b/mindspore/core/ops/not_equal.cc index 60e7f4ce10..e38cfa8587 100644 --- a/mindspore/core/ops/not_equal.cc +++ b/mindspore/core/ops/not_equal.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -36,6 +37,7 @@ TypePtr InferType(const PrimitivePtr &prim, const std::vector & } } // namespace +MIND_API_BASE_IMPL(NotEqual, PrimitiveC, BaseOperator); AbstractBasePtr NotEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/not_equal.h b/mindspore/core/ops/not_equal.h index 75bdbdeb1c..68f59cd274 100644 --- a/mindspore/core/ops/not_equal.h +++ b/mindspore/core/ops/not_equal.h @@ -19,28 +19,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameNotEqual = prim::kNotEqual; +constexpr auto kNameNotEqual = "NotEqual"; /// \brief Computes the non-equivalence of two tensors element-wise. /// Refer to Python API @ref mindspore.ops.NotEqual for more details. -class MS_CORE_API NotEqual : public PrimitiveC { +class MIND_API NotEqual : public BaseOperator { public: + MIND_API_BASE_MEMBER(NotEqual); /// \brief Constructor. - NotEqual() : PrimitiveC(prim::kPrimNotEqual->name()) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~NotEqual() = default; - MS_DECLARE_PARENT(NotEqual, PrimitiveC); + NotEqual() : BaseOperator("NotEqual") { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.NotEqual for the inputs. void Init() const {} }; -AbstractBasePtr NotEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr NotEqualInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/one_hot.cc b/mindspore/core/ops/one_hot.cc index 575b1ee88c..46b98a1442 100644 --- a/mindspore/core/ops/one_hot.cc +++ b/mindspore/core/ops/one_hot.cc @@ -19,11 +19,12 @@ #include "ops/one_hot.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void OneHot::Init(const int64_t axis) { this->set_axis(axis); } -void OneHot::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void OneHot::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t OneHot::get_axis() const { return GetValue(GetAttr(kAxis)); } namespace { @@ -87,6 +88,8 @@ TypePtr OneHotInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/one_hot.h b/mindspore/core/ops/one_hot.h index 1435e90a91..ef1df4a7ad 100644 --- a/mindspore/core/ops/one_hot.h +++ b/mindspore/core/ops/one_hot.h @@ -19,22 +19,17 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Computes a one-hot tensor. Refer to Python API @ref mindspore.ops.OneHot for more details. -class MS_CORE_API OneHot : public PrimitiveC { +class MIND_API OneHot : public BaseOperator { public: + MIND_API_BASE_MEMBER(OneHot); /// \brief Constructor. - OneHot() : PrimitiveC(prim::kPrimOneHot->name()) { - InitIOName({"indices", "depth", "on_value", "off_value"}, {"output"}); - } - /// \brief Destructor. - ~OneHot() = default; - MS_DECLARE_PARENT(OneHot, PrimitiveC); + OneHot() : BaseOperator("OneHot") { InitIOName({"indices", "depth", "on_value", "off_value"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.OneHot for the inputs. void Init(const int64_t axis); /// \brief Set axis. diff --git a/mindspore/core/ops/ones.cc b/mindspore/core/ops/ones.cc index 2b4c1ddbe1..30578fc2cf 100644 --- a/mindspore/core/ops/ones.cc +++ b/mindspore/core/ops/ones.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -73,6 +74,8 @@ ValuePtr OnesInferValue(const PrimitivePtr &prim, const std::vector #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Creates a tensor filled with value ones. Refer to Python API @ref mindspore.ops.Ones for more details. -class Ones : public PrimitiveC { +class Ones : public BaseOperator { public: + MIND_API_BASE_MEMBER(Ones); /// \brief Constructor. - Ones() : PrimitiveC(prim::kPrimOnes->name()) {} - /// \brief Destructor. - ~Ones() = default; - MS_DECLARE_PARENT(Ones, PrimitiveC); + Ones() : BaseOperator("Ones") {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Ones for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/ones_like.cc b/mindspore/core/ops/ones_like.cc index 38a1e75203..71c097f0c3 100644 --- a/mindspore/core/ops/ones_like.cc +++ b/mindspore/core/ops/ones_like.cc @@ -23,6 +23,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -42,6 +43,7 @@ TypePtr OnesLikeInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/ones_like.h b/mindspore/core/ops/ones_like.h index 4edf6aa17c..92a31e0fa2 100644 --- a/mindspore/core/ops/ones_like.h +++ b/mindspore/core/ops/ones_like.h @@ -18,26 +18,25 @@ #define MINDSPORE_CORE_OPS_ONES_LIKE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { +constexpr auto kNameOnesLike = "OnesLike"; /// \brief Creates a new tensor. The values of all elements are 1. /// Refer to Python API @ref mindspore.ops.OnesLike for more details. -class MS_CORE_API OnesLike : public PrimitiveC { +class MIND_API OnesLike : public BaseOperator { public: + MIND_API_BASE_MEMBER(OnesLike); /// \brief Constructor. - OnesLike() : PrimitiveC(prim::kPrimOnesLike->name()) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~OnesLike() = default; - MS_DECLARE_PARENT(OnesLike, PrimitiveC); + OnesLike() : BaseOperator(kNameOnesLike) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.OnesLike for the inputs. void Init() const {} }; -AbstractBasePtr OnesLikeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr OnesLikeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/op_name.h b/mindspore/core/ops/op_name.h new file mode 100644 index 0000000000..0951d3c8f5 --- /dev/null +++ b/mindspore/core/ops/op_name.h @@ -0,0 +1,282 @@ +/** + * Copyright 2022 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. + */ + +#ifndef MINDSPORE_CORE_OPS_OP_NAME_H +#define MINDSPORE_CORE_OPS_OP_NAME_H +#include + +namespace mindspore::ops { +constexpr auto kAlpha = "alpha"; +constexpr auto kActivation = "activation"; +constexpr auto kActivationType = "activation_type"; +constexpr auto kAttentionQActType = "attention_q_act_type"; +constexpr auto kAttentionKActType = "attention_k_act_type"; +constexpr auto kAttentionVActType = "attention_v_act_type"; +constexpr auto kAddress = "address"; +constexpr auto kAlignCorners = "align_corners"; +constexpr auto kAttr = "attr"; +constexpr auto kAspectRatios = "aspect_ratios"; +constexpr auto kAxes = "axes"; +constexpr auto kAxis = "axis"; +constexpr auto kAxisType = "axis_type"; +constexpr auto kBaseSize = "base_size"; +constexpr auto kBatchDim = "batch_dim"; +constexpr auto kBeginMask = "begin_mask"; +constexpr auto kBeginNormAxis = "begin_norm_axis"; +constexpr auto kBeginParamsAxis = "begin_params_axis"; +constexpr auto kBeta = "beta"; +constexpr auto kBias = "bias"; +constexpr auto kBidirectional = "bidirectional"; +constexpr auto kBlockSize = "block_size"; +constexpr auto kBlockShape = "block_shape"; +constexpr auto kCellClip = "cell_clip"; +constexpr auto kCellDepth = "cell_depth"; +constexpr auto kCenterPointBox = "center_point_box"; +constexpr auto kClip = "clip"; +constexpr auto kCondition = "condition"; +constexpr auto kCrops = "crops"; +constexpr auto kEquation = "equation"; +constexpr auto kCustom = "custom"; +constexpr auto kDampening = "dampening"; +constexpr auto kDataType = "data_type"; +constexpr auto kDctCoeffNum = "dct_coeff_num"; +constexpr auto kDelta = "delta"; +constexpr auto kDependMode = "depend_mode"; +constexpr auto kDepthRadius = "depth_radius"; +constexpr auto kDetectionsPerClass = "detections_per_class"; +constexpr auto kDilation = "dilation"; +constexpr auto kDropout = "dropout"; +constexpr auto kDstT = "dst_t"; +constexpr auto kDType = "d_type"; +constexpr auto kEllipsisMask = "ellipsis_mask"; +constexpr auto kEndMask = "end_mask"; +constexpr auto kEps = "eps"; +constexpr auto kEpsilon = "epsilon"; +constexpr auto kElement_dtype = "element_dtype"; +constexpr auto kFeatStride = "feat_stride"; +constexpr auto kFftLength = "fft_length"; +constexpr auto kFilterBankChannelNum = "filter_bank_channel_num"; +constexpr auto kFlip = "flip"; +constexpr auto kFormat = "format"; +constexpr auto kOriginalFormat = "OriginalFormat"; +constexpr auto kFreqLowerLimit = "freq_lower_limit"; +constexpr auto kFreqUpperLimit = "freq_upper_limit"; +constexpr auto kFreezeBn = "freeze_bn"; +constexpr auto kGateOrder = "gate_order"; +constexpr auto kGlobal = "global"; +constexpr auto kGrad = "grad"; +constexpr auto kIsGrad = "is_grad"; +constexpr auto kGradientScale = "gradient_scale"; +constexpr auto kGradX = "grad_x"; +constexpr auto kGradY = "grad_y"; +constexpr auto kGroup = "group"; +constexpr auto kHasBias = "has_bias"; +constexpr auto kAttentionHasMask = "attention_has_mask"; +constexpr auto kHiddenSize = "hidden_size"; +constexpr auto kId = "id"; +constexpr auto kImageSizeH = "image_size_h"; +constexpr auto kImageSizeW = "image_size_w"; +constexpr auto kIncludeALLGrams = "include_all_grams"; +constexpr auto kInputSize = "input_size"; +constexpr auto kInChannel = "in_channel"; +constexpr auto kInputShape = "input_shape"; +constexpr auto kIoFormat = "io_format"; +constexpr auto kIsScale = "is_scale"; +constexpr auto kIsTraining = "is_training"; +constexpr auto kKeepDims = "keep_dims"; +constexpr auto kKeepProb = "keep_prob"; +constexpr auto kKernelSize = "kernel_size"; +constexpr auto kLimit = "limit"; +constexpr auto kMagSquare = "mag_square"; +constexpr auto kMax = "max"; +constexpr auto kMaxSizes = "max_sizes"; +constexpr auto kMaxSkipSize = "max_skip_size"; +constexpr auto kMaxClassesPerDetection = "max_classes_per_detection"; +constexpr auto kMaxDetections = "max_detections"; +constexpr auto kMaxNorm = "max_norm"; +constexpr auto kMin = "min"; +constexpr auto kMinSize = "min_size"; +constexpr auto kMinSizes = "min_sizes"; +constexpr auto kMode = "mode"; +constexpr auto kMomentum = "momentum"; +constexpr auto kN = "n"; +constexpr auto kNarrowRange = "narrow_range"; +constexpr auto kNesterov = "nesterov"; +constexpr auto kNewAxisMask = "new_axis_mask"; +constexpr auto kNgramSize = "ngram_size"; +constexpr auto kNmsThresh = "nms_thresh"; +constexpr auto kNormRegion = "norm_region"; +constexpr auto kNumLayers = "num_layers"; +constexpr auto kNumElements = "num_elements"; +constexpr auto kNumBits = "num_bits"; +constexpr auto kNumDirections = "num_directions"; +constexpr auto kNumProj = "num_proj"; +constexpr auto kAttentionNumHeads = "attention_num_heads"; +constexpr auto kAttentionSizePerHead = "attention_size_per_head"; +constexpr auto kAttentionFromSeqLen = "attention_from_seq_len"; +constexpr auto kAttentionToSeqLen = "attention_to_seq_len"; +constexpr auto kOffset = "offset"; +constexpr auto kNmsIouThreshold = "nms_iou_threshold"; +constexpr auto kNmsScoreThreshold = "nms_score_threshold"; +constexpr auto kNumClasses = "num_classes"; +constexpr auto kOffsets = "offsets"; +constexpr auto kOffsetA = "offset_a"; +constexpr auto kOrder = "order"; +constexpr auto kOutChannel = "out_channel"; +constexpr auto kOutMaxValue = "out_max_value"; +constexpr auto kOutputChannel = "output_channel"; +constexpr auto kOutputNum = "output_num"; +constexpr auto kOutputPaddings = "output_paddings"; +constexpr auto kOutputType = "output_type"; +constexpr auto kOutQuantized = "out_quantized"; +constexpr auto kP = "p"; +constexpr auto kPad = "pad"; +constexpr auto kPadding = "padding"; +constexpr auto kPaddingsElementSize = "paddings_element_size"; +constexpr auto kPaddingsSize = "paddings_size"; +constexpr auto kPadItem = "pad_item"; +constexpr auto kPadList = "pad_list"; +constexpr auto kPadMode = "pad_mode"; +constexpr auto kPads = "pads"; +constexpr auto kPadSize = "pad_size"; +constexpr auto kPooledH = "pooled_h"; +constexpr auto kPooledW = "pooled_w"; +constexpr auto kPoolMode = "pool_mode"; +constexpr auto kCeilMode = "ceil_mode"; +constexpr auto kCountIncludePad = "count_include_pad"; +constexpr auto kDivisorOverride = "divisor_override"; +constexpr auto kPostNmsTopn = "post_nms_topn"; +constexpr auto kPower = "power"; +constexpr auto kPreNmsTopn = "pre_nms_topn"; +constexpr auto kRankSize = "rank_size"; +constexpr auto kRatio = "ratio"; +constexpr auto kReduction = "reduction"; +constexpr auto kRootRank = "root_rank"; +constexpr auto kRoundMode = "round_mode"; +constexpr auto kSame = "same"; +constexpr auto kScale = "scale"; +constexpr auto kSeed = "seed"; +constexpr auto kSeed2 = "seed2"; +constexpr auto kSeqDim = "seq_dim"; +constexpr auto kSetattrFlag = "setattr_flag"; +constexpr auto kShape = "shape"; +constexpr auto kShapeGamma = "shape_gamma"; +constexpr auto kShapeSize = "shape_size"; +constexpr auto kShift = "shift"; +constexpr auto kShrinkAxisMask = "shrink_axis_mask"; +constexpr auto kSize = "size"; +constexpr auto kSorted = "sorted"; +constexpr auto kSrcT = "src_t"; +constexpr auto kStart = "start"; +constexpr auto kStepH = "step_h"; +constexpr auto kStepW = "step_w"; +constexpr auto kStride = "stride"; +constexpr auto kStrides = "strides"; +constexpr auto kShapeType = "shape_type"; +constexpr auto kSubGraphIndex = "sub_graph_index"; +constexpr auto kSummarize = "summarize"; +constexpr auto kTimeMajor = "time_major"; +constexpr auto kTopK = "top_k"; +constexpr auto kTransposeA = "transpose_a"; +constexpr auto kTransposeB = "transpose_b"; +constexpr auto kNegativeSlope = "negative_slope"; +constexpr auto kType = "type"; +constexpr auto kUseAxis = "use_axis"; +constexpr auto kUseLocking = "use_locking"; +constexpr auto kUseNesterov = "use_nesterov"; +constexpr auto kUseNesteroy = "use_nesteroy"; +constexpr auto kUseRegularNms = "use_regular_nms"; +constexpr auto kValid = "valid"; +constexpr auto kValue = "value"; +constexpr auto kVariances = "variances"; +constexpr auto kWeightDecay = "weight_decay"; +constexpr auto kWeightThreshold = "weight_threshold"; +constexpr auto kWindow = "window"; +constexpr auto kWindowSize = "window_size"; +constexpr auto kPaddings = "paddings"; +constexpr auto kInput_size = "input_size"; +constexpr auto kHidden_size = "hidden_size"; +constexpr auto kChannelShared = "channel_shared"; +constexpr auto kSlope = "slope"; +constexpr auto kBase = "base"; +constexpr auto kConstantValue = "constant_value"; +constexpr auto kSizeSplits = "size_splits"; +constexpr auto kDims = "dims"; +constexpr auto kPaddingMode = "padding_mode"; +constexpr auto kLargest = "largest"; +constexpr auto kElementwiseAffine = "elementwise_affine"; +constexpr auto kMinVal = "min_val"; +constexpr auto kMaxVal = "max_val"; +constexpr auto kMethod = "method"; +constexpr auto kNewHeight = "new_height"; +constexpr auto kNewWidth = "new_width"; +constexpr auto kPreserveAspectRatio = "preserve_aspect_ratio"; +constexpr auto kCoordinateTransformMode = "coordinate_transform_mode"; +constexpr auto kCubicCoeff = "cubic_coeff"; +constexpr auto kExcludeOutside = "exclude_outside"; +constexpr auto kExtrapolationValue = "extrapolation_value"; +constexpr auto kNearestMode = "nearest_mode"; +constexpr auto kReduceToEnd = "reduce_to_end"; +constexpr auto kResetAfter = "reset_after"; +constexpr auto kCoeff = "coeff"; +constexpr auto kIsDepthWise = "is_depth_wise"; +constexpr auto kZoneoutCell = "zoneout_cell"; +constexpr auto kZoneoutHidden = "zoneout_hidden"; +constexpr auto kSpliceContext = "context"; +constexpr auto kSpliceForwardIndexes = "forward_indexes"; +constexpr auto kSpliceOutputDims = "output_dim"; +constexpr auto kSideEffectIO = "side_effect_io"; +constexpr auto kDeviceType = "device_type"; +constexpr auto kExclusive = "exclusive"; +constexpr auto kReverse = "reverse"; +constexpr auto kSplitStride = "split_stride"; +constexpr auto kExtendTop = "extend_top"; +constexpr auto kExtendBottom = "extend_bottom"; +constexpr auto kNumberSplit = "number_split"; +constexpr auto kSplitDim = "split_dim"; +constexpr auto kPadTop = "pad_top"; +constexpr auto kTransFormat = "trans_format"; +constexpr auto kApproximate = "approximate"; +constexpr auto kNumOutput = "num_output"; +constexpr auto kUseGlobalStats = "use_global_stats"; +constexpr auto kFmkType = "fmk_type"; +constexpr auto kIsOriginalPadMode = "is_original_pad_mode"; +constexpr auto kOriginalOpName = "original_op_name"; +constexpr auto kSymmetric = "symmetric"; +constexpr auto kDstType = "dst_type"; +constexpr auto kMean = "mean"; + +enum Index : size_t { + kInputIndex0 = 0, + kInputIndex1, + kInputIndex2, + kInputIndex3, + kInputIndex4, + kInputIndex5, + kInputIndex6, + kInputIndex7, + kInputIndex8, + kInputIndex9, + kInputIndex10, + kInputIndex11, + kInputIndex12, + kInputIndex13, + kInputIndex14, + kInputIndex15, + kInputIndex16, +}; +} // namespace mindspore::ops +#endif // MINDSPORE_CORE_OPS_OP_NAME_H diff --git a/mindspore/core/ops/op_utils.cc b/mindspore/core/ops/op_utils.cc index 7a8ded75bb..6be8655857 100644 --- a/mindspore/core/ops/op_utils.cc +++ b/mindspore/core/ops/op_utils.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { diff --git a/mindspore/core/ops/op_utils.h b/mindspore/core/ops/op_utils.h index 377a7198c1..f1e6fe3f3f 100644 --- a/mindspore/core/ops/op_utils.h +++ b/mindspore/core/ops/op_utils.h @@ -22,269 +22,10 @@ #include #include #include "abstract/primitive_infer_map.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/shared_ptr.h" +#include "./op_name.h" namespace mindspore::ops { -constexpr auto kAlpha = "alpha"; -constexpr auto kActivation = "activation"; -constexpr auto kActivationType = "activation_type"; -constexpr auto kAttentionQActType = "attention_q_act_type"; -constexpr auto kAttentionKActType = "attention_k_act_type"; -constexpr auto kAttentionVActType = "attention_v_act_type"; -constexpr auto kAddress = "address"; -constexpr auto kAlignCorners = "align_corners"; -constexpr auto kAttr = "attr"; -constexpr auto kAspectRatios = "aspect_ratios"; -constexpr auto kAxes = "axes"; -constexpr auto kAxis = "axis"; -constexpr auto kAxisType = "axis_type"; -constexpr auto kBaseSize = "base_size"; -constexpr auto kBatchDim = "batch_dim"; -constexpr auto kBeginMask = "begin_mask"; -constexpr auto kBeginNormAxis = "begin_norm_axis"; -constexpr auto kBeginParamsAxis = "begin_params_axis"; -constexpr auto kBeta = "beta"; -constexpr auto kBias = "bias"; -constexpr auto kBidirectional = "bidirectional"; -constexpr auto kBlockSize = "block_size"; -constexpr auto kBlockShape = "block_shape"; -constexpr auto kCellClip = "cell_clip"; -constexpr auto kCellDepth = "cell_depth"; -constexpr auto kCenterPointBox = "center_point_box"; -constexpr auto kClip = "clip"; -constexpr auto kCondition = "condition"; -constexpr auto kCrops = "crops"; -constexpr auto kEquation = "equation"; -constexpr auto kCustom = "custom"; -constexpr auto kDampening = "dampening"; -constexpr auto kDataType = "data_type"; -constexpr auto kDctCoeffNum = "dct_coeff_num"; -constexpr auto kDelta = "delta"; -constexpr auto kDependMode = "depend_mode"; -constexpr auto kDepthRadius = "depth_radius"; -constexpr auto kDetectionsPerClass = "detections_per_class"; -constexpr auto kDilation = "dilation"; -constexpr auto kDropout = "dropout"; -constexpr auto kDstT = "dst_t"; -constexpr auto kDType = "d_type"; -constexpr auto kEllipsisMask = "ellipsis_mask"; -constexpr auto kEndMask = "end_mask"; -constexpr auto kEps = "eps"; -constexpr auto kEpsilon = "epsilon"; -constexpr auto kElement_dtype = "element_dtype"; -constexpr auto kFeatStride = "feat_stride"; -constexpr auto kFftLength = "fft_length"; -constexpr auto kFilterBankChannelNum = "filter_bank_channel_num"; -constexpr auto kFlip = "flip"; -constexpr auto kFormat = "format"; -constexpr auto kOriginalFormat = "OriginalFormat"; -constexpr auto kFreqLowerLimit = "freq_lower_limit"; -constexpr auto kFreqUpperLimit = "freq_upper_limit"; -constexpr auto kFreezeBn = "freeze_bn"; -constexpr auto kGateOrder = "gate_order"; -constexpr auto kGlobal = "global"; -constexpr auto kGrad = "grad"; -constexpr auto kIsGrad = "is_grad"; -constexpr auto kGradientScale = "gradient_scale"; -constexpr auto kGradX = "grad_x"; -constexpr auto kGradY = "grad_y"; -constexpr auto kGroup = "group"; -constexpr auto kHasBias = "has_bias"; -constexpr auto kAttentionHasMask = "attention_has_mask"; -constexpr auto kHiddenSize = "hidden_size"; -constexpr auto kId = "id"; -constexpr auto kImageSizeH = "image_size_h"; -constexpr auto kImageSizeW = "image_size_w"; -constexpr auto kIncludeALLGrams = "include_all_grams"; -constexpr auto kInputSize = "input_size"; -constexpr auto kInChannel = "in_channel"; -constexpr auto kInputShape = "input_shape"; -constexpr auto kIoFormat = "io_format"; -constexpr auto kIsScale = "is_scale"; -constexpr auto kIsTraining = "is_training"; -constexpr auto kKeepDims = "keep_dims"; -constexpr auto kKeepProb = "keep_prob"; -constexpr auto kKernelSize = "kernel_size"; -constexpr auto kLimit = "limit"; -constexpr auto kMagSquare = "mag_square"; -constexpr auto kMax = "max"; -constexpr auto kMaxSizes = "max_sizes"; -constexpr auto kMaxSkipSize = "max_skip_size"; -constexpr auto kMaxClassesPerDetection = "max_classes_per_detection"; -constexpr auto kMaxDetections = "max_detections"; -constexpr auto kMaxNorm = "max_norm"; -constexpr auto kMin = "min"; -constexpr auto kMinSize = "min_size"; -constexpr auto kMinSizes = "min_sizes"; -constexpr auto kMode = "mode"; -constexpr auto kMomentum = "momentum"; -constexpr auto kN = "n"; -constexpr auto kNarrowRange = "narrow_range"; -constexpr auto kNesterov = "nesterov"; -constexpr auto kNewAxisMask = "new_axis_mask"; -constexpr auto kNgramSize = "ngram_size"; -constexpr auto kNmsThresh = "nms_thresh"; -constexpr auto kNormRegion = "norm_region"; -constexpr auto kNumLayers = "num_layers"; -constexpr auto kNumElements = "num_elements"; -constexpr auto kNumBits = "num_bits"; -constexpr auto kNumDirections = "num_directions"; -constexpr auto kNumProj = "num_proj"; -constexpr auto kAttentionNumHeads = "attention_num_heads"; -constexpr auto kAttentionSizePerHead = "attention_size_per_head"; -constexpr auto kAttentionFromSeqLen = "attention_from_seq_len"; -constexpr auto kAttentionToSeqLen = "attention_to_seq_len"; -constexpr auto kOffset = "offset"; -constexpr auto kNmsIouThreshold = "nms_iou_threshold"; -constexpr auto kNmsScoreThreshold = "nms_score_threshold"; -constexpr auto kNumClasses = "num_classes"; -constexpr auto kOffsets = "offsets"; -constexpr auto kOffsetA = "offset_a"; -constexpr auto kOrder = "order"; -constexpr auto kOutChannel = "out_channel"; -constexpr auto kOutMaxValue = "out_max_value"; -constexpr auto kOutputChannel = "output_channel"; -constexpr auto kOutputNum = "output_num"; -constexpr auto kOutputPaddings = "output_paddings"; -constexpr auto kOutputType = "output_type"; -constexpr auto kOutQuantized = "out_quantized"; -constexpr auto kP = "p"; -constexpr auto kPad = "pad"; -constexpr auto kPadding = "padding"; -constexpr auto kPaddingsElementSize = "paddings_element_size"; -constexpr auto kPaddingsSize = "paddings_size"; -constexpr auto kPadItem = "pad_item"; -constexpr auto kPadList = "pad_list"; -constexpr auto kPadMode = "pad_mode"; -constexpr auto kPads = "pads"; -constexpr auto kPadSize = "pad_size"; -constexpr auto kPooledH = "pooled_h"; -constexpr auto kPooledW = "pooled_w"; -constexpr auto kPoolMode = "pool_mode"; -constexpr auto kCeilMode = "ceil_mode"; -constexpr auto kCountIncludePad = "count_include_pad"; -constexpr auto kDivisorOverride = "divisor_override"; -constexpr auto kPostNmsTopn = "post_nms_topn"; -constexpr auto kPower = "power"; -constexpr auto kPreNmsTopn = "pre_nms_topn"; -constexpr auto kRankSize = "rank_size"; -constexpr auto kRatio = "ratio"; -constexpr auto kReduction = "reduction"; -constexpr auto kRootRank = "root_rank"; -constexpr auto kRoundMode = "round_mode"; -constexpr auto kSame = "same"; -constexpr auto kScale = "scale"; -constexpr auto kSeed = "seed"; -constexpr auto kSeed2 = "seed2"; -constexpr auto kSeqDim = "seq_dim"; -constexpr auto kSetattrFlag = "setattr_flag"; -constexpr auto kShape = "shape"; -constexpr auto kShapeGamma = "shape_gamma"; -constexpr auto kShapeSize = "shape_size"; -constexpr auto kShift = "shift"; -constexpr auto kShrinkAxisMask = "shrink_axis_mask"; -constexpr auto kSize = "size"; -constexpr auto kSorted = "sorted"; -constexpr auto kSrcT = "src_t"; -constexpr auto kStart = "start"; -constexpr auto kStepH = "step_h"; -constexpr auto kStepW = "step_w"; -constexpr auto kStride = "stride"; -constexpr auto kStrides = "strides"; -constexpr auto kShapeType = "shape_type"; -constexpr auto kSubGraphIndex = "sub_graph_index"; -constexpr auto kSummarize = "summarize"; -constexpr auto kTimeMajor = "time_major"; -constexpr auto kTopK = "top_k"; -constexpr auto kTransposeA = "transpose_a"; -constexpr auto kTransposeB = "transpose_b"; -constexpr auto kNegativeSlope = "negative_slope"; -constexpr auto kType = "type"; -constexpr auto kUseAxis = "use_axis"; -constexpr auto kUseLocking = "use_locking"; -constexpr auto kUseNesterov = "use_nesterov"; -constexpr auto kUseNesteroy = "use_nesteroy"; -constexpr auto kUseRegularNms = "use_regular_nms"; -constexpr auto kValid = "valid"; -constexpr auto kValue = "value"; -constexpr auto kVariances = "variances"; -constexpr auto kWeightDecay = "weight_decay"; -constexpr auto kWeightThreshold = "weight_threshold"; -constexpr auto kWindow = "window"; -constexpr auto kWindowSize = "window_size"; -constexpr auto kPaddings = "paddings"; -constexpr auto kInput_size = "input_size"; -constexpr auto kHidden_size = "hidden_size"; -constexpr auto kChannelShared = "channel_shared"; -constexpr auto kSlope = "slope"; -constexpr auto kBase = "base"; -constexpr auto kConstantValue = "constant_value"; -constexpr auto kSizeSplits = "size_splits"; -constexpr auto kDims = "dims"; -constexpr auto kPaddingMode = "padding_mode"; -constexpr auto kLargest = "largest"; -constexpr auto kElementwiseAffine = "elementwise_affine"; -constexpr auto kMinVal = "min_val"; -constexpr auto kMaxVal = "max_val"; -constexpr auto kMethod = "method"; -constexpr auto kNewHeight = "new_height"; -constexpr auto kNewWidth = "new_width"; -constexpr auto kPreserveAspectRatio = "preserve_aspect_ratio"; -constexpr auto kCoordinateTransformMode = "coordinate_transform_mode"; -constexpr auto kCubicCoeff = "cubic_coeff"; -constexpr auto kExcludeOutside = "exclude_outside"; -constexpr auto kExtrapolationValue = "extrapolation_value"; -constexpr auto kNearestMode = "nearest_mode"; -constexpr auto kReduceToEnd = "reduce_to_end"; -constexpr auto kResetAfter = "reset_after"; -constexpr auto kCoeff = "coeff"; -constexpr auto kIsDepthWise = "is_depth_wise"; -constexpr auto kZoneoutCell = "zoneout_cell"; -constexpr auto kZoneoutHidden = "zoneout_hidden"; -constexpr auto kSpliceContext = "context"; -constexpr auto kSpliceForwardIndexes = "forward_indexes"; -constexpr auto kSpliceOutputDims = "output_dim"; -constexpr auto kSideEffectIO = "side_effect_io"; -constexpr auto kDeviceType = "device_type"; -constexpr auto kExclusive = "exclusive"; -constexpr auto kReverse = "reverse"; -constexpr auto kSplitStride = "split_stride"; -constexpr auto kExtendTop = "extend_top"; -constexpr auto kExtendBottom = "extend_bottom"; -constexpr auto kNumberSplit = "number_split"; -constexpr auto kSplitDim = "split_dim"; -constexpr auto kPadTop = "pad_top"; -constexpr auto kTransFormat = "trans_format"; -constexpr auto kApproximate = "approximate"; -constexpr auto kNumOutput = "num_output"; -constexpr auto kUseGlobalStats = "use_global_stats"; -constexpr auto kFmkType = "fmk_type"; -constexpr auto kIsOriginalPadMode = "is_original_pad_mode"; -constexpr auto kOriginalOpName = "original_op_name"; -constexpr auto kSymmetric = "symmetric"; -constexpr auto kDstType = "dst_type"; -constexpr auto kMean = "mean"; - -enum Index : size_t { - kInputIndex0 = 0, - kInputIndex1, - kInputIndex2, - kInputIndex3, - kInputIndex4, - kInputIndex5, - kInputIndex6, - kInputIndex7, - kInputIndex8, - kInputIndex9, - kInputIndex10, - kInputIndex11, - kInputIndex12, - kInputIndex13, - kInputIndex14, - kInputIndex15, - kInputIndex16, -}; - const std::set common_valid_types = {kInt8, kInt16, kInt32, kInt64, kUInt8, kUInt16, kUInt32, kUInt64, kFloat16, kFloat32, kFloat64}; @@ -303,6 +44,16 @@ const std::set all_types = { std::vector CalBroadCastShape(std::vector x_shape, std::vector y_shape, const std::string &op_name, const std::string &op_x_name = "input1", const std::string &op_y_name = "input2"); -abstract::ShapePtr BroadCastInferShape(const std::string &op_name, const std::vector &input_args); +abstract::ShapePtr BroadCastInferShape(const std::string &op_name, + const std::vector &input_args); + +template +api::SharedPtr GetOperator(const AnfNodePtr &node) { + auto prim = GetValueNode(node); + if (prim == nullptr) { + return nullptr; + } + return api::MakeShared(prim); +} } // namespace mindspore::ops #endif // MINDSPORE_CORE_OPS_OP_UTILS_H diff --git a/mindspore/core/ops/pack.cc b/mindspore/core/ops/pack.cc index 6ddb26ca16..357f191c9a 100644 --- a/mindspore/core/ops/pack.cc +++ b/mindspore/core/ops/pack.cc @@ -15,15 +15,20 @@ */ #include "ops/pack.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Pack::set_axis(const int64_t &axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void Pack::set_axis(const int64_t &axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } int64_t Pack::get_axis() const { return GetValue(GetAttr(kAxis)); } void Pack::Init(const int64_t &axis) { this->set_axis(axis); } +MIND_API_BASE_IMPL(Pack, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNamePack, Pack); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/pack.h b/mindspore/core/ops/pack.h index 44a6da5eb4..3411e58064 100644 --- a/mindspore/core/ops/pack.h +++ b/mindspore/core/ops/pack.h @@ -22,24 +22,20 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePack = "Pack"; /// \brief Stacks a list of tensors in specified axis. /// Refer to Python API @ref mindspore.ops.Stack for more details. -class MS_CORE_API Pack : public PrimitiveC { +class MIND_API Pack : public BaseOperator { public: + MIND_API_BASE_MEMBER(Pack); /// \brief Constructor. - Pack() : PrimitiveC(kNamePack) {} - /// \brief Destructor. - ~Pack() = default; - MS_DECLARE_PARENT(Pack, PrimitiveC); + Pack() : BaseOperator(kNamePack) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Stack for the inputs. void Init(const int64_t &axis = 0); /// \brief Set axis. @@ -49,8 +45,8 @@ class MS_CORE_API Pack : public PrimitiveC { /// \return axis. int64_t get_axis() const; }; -AbstractBasePtr PackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr PackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_PACK_H_ diff --git a/mindspore/core/ops/pad.cc b/mindspore/core/ops/pad.cc index a457a3f17a..d091226308 100644 --- a/mindspore/core/ops/pad.cc +++ b/mindspore/core/ops/pad.cc @@ -17,17 +17,20 @@ #include #include "ops/pad.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void Pad::Init(const std::vector> &paddings) { this->set_paddings(paddings); } void Pad::set_paddings(const std::vector> &paddings) { - (void)this->AddAttr(kPaddings, MakeValue(paddings)); + (void)this->AddAttr(kPaddings, api::MakeValue(paddings)); } std::vector> Pad::get_paddings() const { return GetValue>>(GetAttr(kPaddings)); } +MIND_API_BASE_IMPL(Pad, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNamePad, Pad); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/pad.h b/mindspore/core/ops/pad.h index ef1769251d..90f339b232 100644 --- a/mindspore/core/ops/pad.h +++ b/mindspore/core/ops/pad.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePad = "Pad"; /// \brief Pads the input tensor according to the paddings. Refer to Python API @ref mindspore.ops.Pad for more details. -class MS_CORE_API Pad : public PrimitiveC { +class MIND_API Pad : public BaseOperator { public: + MIND_API_BASE_MEMBER(Pad); /// \brief Constructor. - Pad() : PrimitiveC(kNamePad) { InitIOName({"x"}, {"y"}); } - explicit Pad(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Pad() = default; - MS_DECLARE_PARENT(Pad, PrimitiveC); + Pad() : BaseOperator(kNamePad) { InitIOName({"x"}, {"y"}); } + explicit Pad(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Pad for the inputs. void Init(const std::vector> &paddings); /// \brief Set paddings. @@ -45,8 +43,8 @@ class MS_CORE_API Pad : public PrimitiveC { /// \return paddings. std::vector> get_paddings() const; }; -AbstractBasePtr PadInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr PadInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/partial.cc b/mindspore/core/ops/partial.cc index 8105873206..71b7866a80 100644 --- a/mindspore/core/ops/partial.cc +++ b/mindspore/core/ops/partial.cc @@ -16,9 +16,11 @@ #include "ops/partial.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Partial, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNamePartial, Partial); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/partial.h b/mindspore/core/ops/partial.h index f7f636e513..cd798cc6a9 100644 --- a/mindspore/core/ops/partial.h +++ b/mindspore/core/ops/partial.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_PARTIAL_H_ #define MINDSPORE_CORE_OPS_PARTIAL_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePartial = "Partial"; /// \brief Partial defined Partial operator prototype of lite. -class MS_CORE_API Partial : public PrimitiveC { +class MIND_API Partial : public BaseOperator { public: + MIND_API_BASE_MEMBER(Partial); /// \brief Constructor. - Partial() : PrimitiveC(kNamePartial) {} - - /// \brief Destructor. - ~Partial() = default; - - MS_DECLARE_PARENT(Partial, PrimitiveC); + Partial() : BaseOperator(kNamePartial) {} /// \brief Method to init the op's attributes. void Init() const {} diff --git a/mindspore/core/ops/pow.cc b/mindspore/core/ops/pow.cc index 2bad972f3b..dbd3e8653e 100644 --- a/mindspore/core/ops/pow.cc +++ b/mindspore/core/ops/pow.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -74,6 +75,8 @@ TypePtr PowInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto prim_name = primitive->name(); diff --git a/mindspore/core/ops/pow.h b/mindspore/core/ops/pow.h index a00f324f56..5c85493296 100644 --- a/mindspore/core/ops/pow.h +++ b/mindspore/core/ops/pow.h @@ -20,28 +20,25 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePow = "Pow"; /// \brief Computes a tensor to the power of the second input. /// Refer to Python API @ref mindspore.ops.Pow for more details. -class MS_CORE_API Pow : public PrimitiveC { +class MIND_API Pow : public BaseOperator { public: + MIND_API_BASE_MEMBER(Pow); /// \brief Constructor. - explicit Pow(const std::string &k_name = kNamePow) : PrimitiveC(k_name) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Pow() = default; - MS_DECLARE_PARENT(Pow, PrimitiveC); + explicit Pow(const std::string &k_name = kNamePow) : BaseOperator(k_name) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Pow for the inputs. void Init(); }; -AbstractBasePtr PowInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr PowInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimPowPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/prelu.cc b/mindspore/core/ops/prelu.cc index 4be4d00454..0672c7a15d 100644 --- a/mindspore/core/ops/prelu.cc +++ b/mindspore/core/ops/prelu.cc @@ -14,11 +14,16 @@ * limitations under the License. */ -#include #include "ops/prelu.h" +#include +#include +#include "ops/primitive_c.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PReLU, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNamePReLU, PReLU); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/prelu.h b/mindspore/core/ops/prelu.h index 024822fed0..2147a21bb3 100644 --- a/mindspore/core/ops/prelu.h +++ b/mindspore/core/ops/prelu.h @@ -19,29 +19,27 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePReLU = "PReLU"; /// \brief Parametric Rectified Linear Unit activation function. /// Refer to Python API @ref mindspore.ops.PReLU for more details. -class MS_CORE_API PReLU : public PrimitiveC { +class MIND_API PReLU : public BaseOperator { public: + MIND_API_BASE_MEMBER(PReLU); /// \brief Constructor. - PReLU() : PrimitiveC(kNamePReLU) { InitIOName({"x"}, {"y"}); } - explicit PReLU(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~PReLU() = default; - MS_DECLARE_PARENT(PReLU, PrimitiveC); + PReLU() : BaseOperator(kNamePReLU) { InitIOName({"x"}, {"y"}); } + explicit PReLU(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.PReLU for the inputs. void Init() const {} }; -AbstractBasePtr PReLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr PReLUInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/primitive_c.h b/mindspore/core/ops/primitive_c.h index cd6f2acfeb..ae00ffe2e4 100644 --- a/mindspore/core/ops/primitive_c.h +++ b/mindspore/core/ops/primitive_c.h @@ -93,11 +93,11 @@ class MS_CORE_API OpPrimCRegisterHelper { int id_{0}; }; -#define REGISTER_PRIMITIVE_C(kname, primc) \ - std::shared_ptr GetDefaultPrimC##primc() { \ - auto out = std::make_shared(); \ - return out; \ - } \ +#define REGISTER_PRIMITIVE_C(kname, primc) \ + std::shared_ptr GetDefaultPrimC##primc() { \ + primc out; \ + return std::dynamic_pointer_cast(out.impl()); \ + } \ OpPrimCRegisterHelper primc_gen_##kname(kname, GetDefaultPrimC##primc); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/prior_box.cc b/mindspore/core/ops/prior_box.cc index c3b45c539b..86d3dd11f8 100644 --- a/mindspore/core/ops/prior_box.cc +++ b/mindspore/core/ops/prior_box.cc @@ -19,17 +19,19 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(PriorBox, PrimitiveC, BaseOperator); void PriorBox::set_min_sizes(const std::vector &min_sizes) { - (void)this->AddAttr(kMinSizes, MakeValue(min_sizes)); + (void)this->AddAttr(kMinSizes, api::MakeValue(min_sizes)); } std::vector PriorBox::get_min_sizes() const { return GetValue>(GetAttr(kMinSizes)); } void PriorBox::set_max_sizes(const std::vector &max_sizes) { - (void)this->AddAttr(kMaxSizes, MakeValue(max_sizes)); + (void)this->AddAttr(kMaxSizes, api::MakeValue(max_sizes)); } std::vector PriorBox::get_max_sizes() const { @@ -38,13 +40,13 @@ std::vector PriorBox::get_max_sizes() const { } void PriorBox::set_aspect_ratios(const std::vector &aspect_ratios) { - (void)this->AddAttr(kAspectRatios, MakeValue(aspect_ratios)); + (void)this->AddAttr(kAspectRatios, api::MakeValue(aspect_ratios)); } std::vector PriorBox::get_aspect_ratios() const { return GetValue>(GetAttr(kAspectRatios)); } void PriorBox::set_variances(const std::vector &variances) { - (void)this->AddAttr(kVariances, MakeValue(variances)); + (void)this->AddAttr(kVariances, api::MakeValue(variances)); } std::vector PriorBox::get_variances() const { @@ -53,7 +55,7 @@ std::vector PriorBox::get_variances() const { } void PriorBox::set_image_size_w(const int64_t image_size_w) { - (void)this->AddAttr(kImageSizeW, MakeValue(image_size_w)); + (void)this->AddAttr(kImageSizeW, api::MakeValue(image_size_w)); } int64_t PriorBox::get_image_size_w() const { @@ -62,7 +64,7 @@ int64_t PriorBox::get_image_size_w() const { } void PriorBox::set_image_size_h(const int64_t image_size_h) { - (void)this->AddAttr(kImageSizeH, MakeValue(image_size_h)); + (void)this->AddAttr(kImageSizeH, api::MakeValue(image_size_h)); } int64_t PriorBox::get_image_size_h() const { @@ -70,32 +72,32 @@ int64_t PriorBox::get_image_size_h() const { return GetValue(value_ptr); } -void PriorBox::set_step_w(const float step_w) { (void)this->AddAttr(kStepW, MakeValue(step_w)); } +void PriorBox::set_step_w(const float step_w) { (void)this->AddAttr(kStepW, api::MakeValue(step_w)); } float PriorBox::get_step_w() const { auto value_ptr = GetAttr(kStepW); return GetValue(value_ptr); } -void PriorBox::set_step_h(const float step_h) { (void)this->AddAttr(kStepH, MakeValue(step_h)); } +void PriorBox::set_step_h(const float step_h) { (void)this->AddAttr(kStepH, api::MakeValue(step_h)); } float PriorBox::get_step_h() const { auto value_ptr = GetAttr(kStepH); return GetValue(value_ptr); } -void PriorBox::set_clip(const bool clip) { (void)this->AddAttr(kClip, MakeValue(clip)); } +void PriorBox::set_clip(const bool clip) { (void)this->AddAttr(kClip, api::MakeValue(clip)); } bool PriorBox::get_clip() const { auto value_ptr = GetAttr(kClip); return GetValue(value_ptr); } -void PriorBox::set_flip(const bool flip) { (void)this->AddAttr(kFlip, MakeValue(flip)); } +void PriorBox::set_flip(const bool flip) { (void)this->AddAttr(kFlip, api::MakeValue(flip)); } bool PriorBox::get_flip() const { return GetValue(GetAttr(kFlip)); } -void PriorBox::set_offset(const float offset) { (void)this->AddAttr(kOffset, MakeValue(offset)); } +void PriorBox::set_offset(const float offset) { (void)this->AddAttr(kOffset, api::MakeValue(offset)); } float PriorBox::get_offset() const { auto value_ptr = GetAttr(kOffset); diff --git a/mindspore/core/ops/prior_box.h b/mindspore/core/ops/prior_box.h index ed55bed2dd..704533b001 100644 --- a/mindspore/core/ops/prior_box.h +++ b/mindspore/core/ops/prior_box.h @@ -19,23 +19,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNamePriorBox = "PriorBox"; /// \brief PriorBox defined PriorBox operator prototype of lite. -class MS_CORE_API PriorBox : public PrimitiveC { +class MIND_API PriorBox : public BaseOperator { public: + MIND_API_BASE_MEMBER(PriorBox); /// \brief Constructor. - PriorBox() : PrimitiveC(kNamePriorBox) {} - - /// \brief Destructor. - ~PriorBox() = default; - - MS_DECLARE_PARENT(PriorBox, PrimitiveC); + PriorBox() : BaseOperator(kNamePriorBox) {} /// \brief Method to init the op's attributes. /// @@ -168,8 +163,8 @@ class MS_CORE_API PriorBox : public PrimitiveC { float get_offset() const; }; -AbstractBasePtr PriorBoxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr PriorBoxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/proposal.cc b/mindspore/core/ops/proposal.cc index 427de07be3..fe7710f212 100644 --- a/mindspore/core/ops/proposal.cc +++ b/mindspore/core/ops/proposal.cc @@ -20,38 +20,42 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Proposal::set_feat_stride(const float feat_stride) { (void)this->AddAttr(kFeatStride, MakeValue(feat_stride)); } +MIND_API_BASE_IMPL(Proposal, PrimitiveC, BaseOperator); +void Proposal::set_feat_stride(const float feat_stride) { + (void)this->AddAttr(kFeatStride, api::MakeValue(feat_stride)); +} float Proposal::get_feat_stride() const { auto value_ptr = GetAttr(kFeatStride); return GetValue(value_ptr); } -void Proposal::set_base_size(const float base_size) { (void)this->AddAttr(kBaseSize, MakeValue(base_size)); } +void Proposal::set_base_size(const float base_size) { (void)this->AddAttr(kBaseSize, api::MakeValue(base_size)); } float Proposal::get_base_size() const { auto value_ptr = GetAttr(kBaseSize); return GetValue(value_ptr); } -void Proposal::set_min_size(const float min_size) { (void)this->AddAttr(kMinSize, MakeValue(min_size)); } +void Proposal::set_min_size(const float min_size) { (void)this->AddAttr(kMinSize, api::MakeValue(min_size)); } float Proposal::get_min_size() const { auto value_ptr = GetAttr(kMinSize); return GetValue(value_ptr); } -void Proposal::set_ratio(const std::vector &ratio) { (void)this->AddAttr(kRatio, MakeValue(ratio)); } +void Proposal::set_ratio(const std::vector &ratio) { (void)this->AddAttr(kRatio, api::MakeValue(ratio)); } std::vector Proposal::get_ratio() const { auto value_ptr = GetAttr(kRatio); return GetValue>(value_ptr); } -void Proposal::set_scale(const std::vector &scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } +void Proposal::set_scale(const std::vector &scale) { (void)this->AddAttr(kScale, api::MakeValue(scale)); } std::vector Proposal::get_scale() const { auto value_ptr = GetAttr(kScale); @@ -59,7 +63,7 @@ std::vector Proposal::get_scale() const { } void Proposal::set_pre_nms_topn(const int64_t pre_nms_topn) { - (void)this->AddAttr(kPreNmsTopn, MakeValue(pre_nms_topn)); + (void)this->AddAttr(kPreNmsTopn, api::MakeValue(pre_nms_topn)); } int64_t Proposal::get_pre_nms_topn() const { @@ -68,7 +72,7 @@ int64_t Proposal::get_pre_nms_topn() const { } void Proposal::set_post_nms_topn(const int64_t post_nms_topn) { - (void)this->AddAttr(kPostNmsTopn, MakeValue(post_nms_topn)); + (void)this->AddAttr(kPostNmsTopn, api::MakeValue(post_nms_topn)); } int64_t Proposal::get_post_nms_topn() const { @@ -76,7 +80,7 @@ int64_t Proposal::get_post_nms_topn() const { return GetValue(value_ptr); } -void Proposal::set_nms_thresh(const float nms_thresh) { (void)this->AddAttr(kNmsThresh, MakeValue(nms_thresh)); } +void Proposal::set_nms_thresh(const float nms_thresh) { (void)this->AddAttr(kNmsThresh, api::MakeValue(nms_thresh)); } float Proposal::get_nms_thresh() const { auto value_ptr = GetAttr(kNmsThresh); diff --git a/mindspore/core/ops/proposal.h b/mindspore/core/ops/proposal.h index e5ae7c2228..aa554ea598 100644 --- a/mindspore/core/ops/proposal.h +++ b/mindspore/core/ops/proposal.h @@ -18,18 +18,16 @@ #define MINDSPORE_CORE_OPS_PROPOSAL_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameProposal = "Proposal"; -class MS_CORE_API Proposal : public PrimitiveC { +class MIND_API Proposal : public BaseOperator { public: - Proposal() : PrimitiveC(kNameProposal) {} - ~Proposal() = default; - MS_DECLARE_PARENT(Proposal, PrimitiveC); + MIND_API_BASE_MEMBER(Proposal); + Proposal() : BaseOperator(kNameProposal) {} void Init(const float feat_stride, const float base_size, const float min_size, const std::vector &ratio, const std::vector &scale, const int64_t pre_nms_topn, const int64_t post_nms_topn, diff --git a/mindspore/core/ops/quant_dtype_cast.cc b/mindspore/core/ops/quant_dtype_cast.cc index 83fcf9ccd6..60646e9773 100644 --- a/mindspore/core/ops/quant_dtype_cast.cc +++ b/mindspore/core/ops/quant_dtype_cast.cc @@ -15,15 +15,20 @@ */ #include "ops/quant_dtype_cast.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void QuantDTypeCast::set_src_t(const int64_t src_t) { (void)AddAttr(kSrcT, MakeValue(src_t)); } +MIND_API_BASE_IMPL(QuantDTypeCast, PrimitiveC, BaseOperator); +void QuantDTypeCast::set_src_t(const int64_t src_t) { (void)AddAttr(kSrcT, api::MakeValue(src_t)); } int64_t QuantDTypeCast::get_src_t() const { auto value_ptr = this->GetAttr(kSrcT); return GetValue(value_ptr); } -void QuantDTypeCast::set_dst_t(const int64_t dst_t) { (void)AddAttr(kDstT, MakeValue(dst_t)); } +void QuantDTypeCast::set_dst_t(const int64_t dst_t) { (void)AddAttr(kDstT, api::MakeValue(dst_t)); } int64_t QuantDTypeCast::get_dst_t() const { return GetValue(GetAttr(kDstT)); } void QuantDTypeCast::Init(const int64_t src_t, const int64_t dst_t) { this->set_src_t(src_t); diff --git a/mindspore/core/ops/quant_dtype_cast.h b/mindspore/core/ops/quant_dtype_cast.h index cf3720ba42..f598331e15 100644 --- a/mindspore/core/ops/quant_dtype_cast.h +++ b/mindspore/core/ops/quant_dtype_cast.h @@ -22,25 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameQuantDTypeCast = "QuantDTypeCast"; /// \brief QuantDTypeCast QuantDTypeCast the QuantDTypeCast operator prototype. -class MS_CORE_API QuantDTypeCast : public PrimitiveC { +class MIND_API QuantDTypeCast : public BaseOperator { public: + MIND_API_BASE_MEMBER(QuantDTypeCast); /// \brief Constructor. - QuantDTypeCast() : PrimitiveC(kNameQuantDTypeCast) {} - - /// \brief Destructor. - ~QuantDTypeCast() = default; - - MS_DECLARE_PARENT(QuantDTypeCast, PrimitiveC); + QuantDTypeCast() : BaseOperator(kNameQuantDTypeCast) {} /// \brief Method to init the op's attributes. /// @@ -68,8 +62,8 @@ class MS_CORE_API QuantDTypeCast : public PrimitiveC { /// \return the data type of output. int64_t get_dst_t() const; }; -AbstractBasePtr QuantDTypeCastInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr QuantDTypeCastInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/ragged_range.cc b/mindspore/core/ops/ragged_range.cc index e1ef4d16b3..e9023877e0 100644 --- a/mindspore/core/ops/ragged_range.cc +++ b/mindspore/core/ops/ragged_range.cc @@ -15,9 +15,11 @@ */ #include "ops/ragged_range.h" #include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(RaggedRange, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameRaggedRange, RaggedRange); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/ragged_range.h b/mindspore/core/ops/ragged_range.h index 01b68a9199..87b7d8f1fa 100644 --- a/mindspore/core/ops/ragged_range.h +++ b/mindspore/core/ops/ragged_range.h @@ -19,21 +19,18 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRaggedRange = "RaggedRange"; /// \brief RaggedRange operator prototype. -class MS_CORE_API RaggedRange : public PrimitiveC { +class MIND_API RaggedRange : public BaseOperator { public: + MIND_API_BASE_MEMBER(RaggedRange); /// \brief Constructor - RaggedRange() : PrimitiveC(kNameRaggedRange) {} - /// \brief Destructor - ~RaggedRange() = default; - MS_DECLARE_PARENT(RaggedRange, PrimitiveC); + RaggedRange() : BaseOperator(kNameRaggedRange) {} /// \brief Method to init the op. void Init() const {} }; diff --git a/mindspore/core/ops/random_normal.cc b/mindspore/core/ops/random_normal.cc index 94c5bfa88e..cfac627a4c 100644 --- a/mindspore/core/ops/random_normal.cc +++ b/mindspore/core/ops/random_normal.cc @@ -16,20 +16,22 @@ #include "ops/random_normal.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(RandomNormal, PrimitiveC, BaseOperator); void RandomNormal::Init(float seed, float mean, float scale) { this->set_seed(seed); this->set_mean(mean); this->set_scale(scale); } -void RandomNormal::set_seed(float seed) { (void)this->AddAttr(kSeed, MakeValue(seed)); } +void RandomNormal::set_seed(float seed) { (void)this->AddAttr(kSeed, api::MakeValue(seed)); } -void RandomNormal::set_mean(float mean) { (void)this->AddAttr(kMean, MakeValue(mean)); } +void RandomNormal::set_mean(float mean) { (void)this->AddAttr(kMean, api::MakeValue(mean)); } -void RandomNormal::set_scale(float scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } +void RandomNormal::set_scale(float scale) { (void)this->AddAttr(kScale, api::MakeValue(scale)); } float RandomNormal::get_seed() const { auto value_ptr = GetAttr(kSeed); diff --git a/mindspore/core/ops/random_normal.h b/mindspore/core/ops/random_normal.h index 0cd649b1a2..ad6ee52bfc 100644 --- a/mindspore/core/ops/random_normal.h +++ b/mindspore/core/ops/random_normal.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_RANDOM_NORMAL_H_ #define MINDSPORE_CORE_OPS_RANDOM_NORMAL_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRandomNormal = "RandomNormal"; /// \brief RandomNormal defined RandomNormal operator prototype of lite. -class MS_CORE_API RandomNormal : public PrimitiveC { +class MIND_API RandomNormal : public BaseOperator { public: + MIND_API_BASE_MEMBER(RandomNormal); /// \brief Constructor. - RandomNormal() : PrimitiveC(kNameRandomNormal) {} - - /// \brief Destructor. - ~RandomNormal() = default; - - MS_DECLARE_PARENT(RandomNormal, PrimitiveC); + RandomNormal() : BaseOperator(kNameRandomNormal) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/random_standard_normal.cc b/mindspore/core/ops/random_standard_normal.cc index b769402969..485dff4977 100644 --- a/mindspore/core/ops/random_standard_normal.cc +++ b/mindspore/core/ops/random_standard_normal.cc @@ -18,17 +18,19 @@ #include #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(RandomStandardNormal, PrimitiveC, BaseOperator); void RandomStandardNormal::Init(const int64_t seed, const int64_t seed2) { this->set_seed(seed); this->set_seed2(seed2); } -void RandomStandardNormal::set_seed(int64_t seed) { (void)this->AddAttr(kSeed, MakeValue(seed)); } +void RandomStandardNormal::set_seed(int64_t seed) { (void)this->AddAttr(kSeed, api::MakeValue(seed)); } -void RandomStandardNormal::set_seed2(int64_t seed2) { (void)this->AddAttr(kSeed2, MakeValue(seed2)); } +void RandomStandardNormal::set_seed2(int64_t seed2) { (void)this->AddAttr(kSeed2, api::MakeValue(seed2)); } int64_t RandomStandardNormal::get_seed() const { auto value_ptr = GetAttr(kSeed); diff --git a/mindspore/core/ops/random_standard_normal.h b/mindspore/core/ops/random_standard_normal.h index 801a62146a..d88c99da41 100644 --- a/mindspore/core/ops/random_standard_normal.h +++ b/mindspore/core/ops/random_standard_normal.h @@ -20,23 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRandomStandardNormal = "RandomStandardNormal"; /// \brief RandomStandardNormal defined RandomStandardNormal operator prototype of lite. -class MS_CORE_API RandomStandardNormal : public PrimitiveC { +class MIND_API RandomStandardNormal : public BaseOperator { public: + MIND_API_BASE_MEMBER(RandomStandardNormal); /// \brief Constructor. - RandomStandardNormal() : PrimitiveC(kNameRandomStandardNormal) {} - - /// \brief Destructor. - ~RandomStandardNormal() = default; - - MS_DECLARE_PARENT(RandomStandardNormal, PrimitiveC); + RandomStandardNormal() : BaseOperator(kNameRandomStandardNormal) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/range.cc b/mindspore/core/ops/range.cc index 749ef7950d..965ac8f85d 100644 --- a/mindspore/core/ops/range.cc +++ b/mindspore/core/ops/range.cc @@ -22,28 +22,30 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Range::set_d_type(const int64_t d_type) { (void)this->AddAttr(kDType, MakeValue(d_type)); } +MIND_API_BASE_IMPL(Range, PrimitiveC, BaseOperator); +void Range::set_d_type(const int64_t d_type) { (void)this->AddAttr(kDType, api::MakeValue(d_type)); } int64_t Range::get_d_type() const { auto value_ptr = GetAttr(kDType); return GetValue(value_ptr); } -void Range::set_start(const int64_t start) { (void)this->AddAttr(kStart, MakeValue(start)); } +void Range::set_start(const int64_t start) { (void)this->AddAttr(kStart, api::MakeValue(start)); } int64_t Range::get_start() const { return GetValue(GetAttr(kStart)); } -void Range::set_limit(const int64_t limit) { (void)this->AddAttr(kLimit, MakeValue(limit)); } +void Range::set_limit(const int64_t limit) { (void)this->AddAttr(kLimit, api::MakeValue(limit)); } int64_t Range::get_limit() const { auto value_ptr = GetAttr(kLimit); return GetValue(value_ptr); } -void Range::set_delta(const int64_t delta) { (void)this->AddAttr(kDelta, MakeValue(delta)); } +void Range::set_delta(const int64_t delta) { (void)this->AddAttr(kDelta, api::MakeValue(delta)); } int64_t Range::get_delta() const { auto value_ptr = GetAttr(kDelta); diff --git a/mindspore/core/ops/range.h b/mindspore/core/ops/range.h index 98bf08b6ce..add76e5769 100644 --- a/mindspore/core/ops/range.h +++ b/mindspore/core/ops/range.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRange = "Range"; /// \brief Creates a sequence of numbers in range [start, limit) with step size delta. /// Refer to Python API @ref mindspore.nn.Range for more details. -class MS_CORE_API Range : public PrimitiveC { +class MIND_API Range : public BaseOperator { public: + MIND_API_BASE_MEMBER(Range); /// \brief Constructor. - Range() : PrimitiveC(kNameRange) {} - /// \brief Destructor. - ~Range() = default; - MS_DECLARE_PARENT(Range, PrimitiveC); + Range() : BaseOperator(kNameRange) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.nn.Range for the inputs. void Init(const int64_t d_type, const int64_t start, const int64_t limit, const int64_t delta); /// \brief Set d_type. @@ -64,8 +62,8 @@ class MS_CORE_API Range : public PrimitiveC { int64_t get_delta() const; }; -AbstractBasePtr RangeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RangeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/rank.cc b/mindspore/core/ops/rank.cc index 2f21147632..65286e9726 100644 --- a/mindspore/core/ops/rank.cc +++ b/mindspore/core/ops/rank.cc @@ -15,9 +15,13 @@ */ #include "ops/rank.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Rank, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameRank, Rank); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/rank.h b/mindspore/core/ops/rank.h index 643604cdad..7a1603b7b9 100644 --- a/mindspore/core/ops/rank.h +++ b/mindspore/core/ops/rank.h @@ -19,27 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRank = "Rank"; /// \brief Returns the rank of a tensor. Refer to Python API @ref mindspore.ops.Rank for more details. -class MS_CORE_API Rank : public PrimitiveC { +class MIND_API Rank : public BaseOperator { public: + MIND_API_BASE_MEMBER(Rank); /// \brief Constructor. - Rank() : PrimitiveC(kNameRank) { auto prim_name = name(); } - /// \brief Destructor. - ~Rank() = default; - MS_DECLARE_PARENT(Rank, PrimitiveC); + Rank() : BaseOperator(kNameRank) { auto prim_name = name(); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Rank for the inputs. void Init() const {} }; -AbstractBasePtr RankInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RankInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_RANK_H_ diff --git a/mindspore/core/ops/real.cc b/mindspore/core/ops/real.cc index b6ff6925f7..c2d7ae6dc5 100644 --- a/mindspore/core/ops/real.cc +++ b/mindspore/core/ops/real.cc @@ -20,6 +20,8 @@ #include #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -58,6 +60,8 @@ AbstractBasePtr RealInfer(const abstract::AnalysisEnginePtr &, const PrimitivePt return abstract::MakeAbstract(RealInferShape(primitive, input_args), RealInferType(primitive, input_args)); } } // namespace + +MIND_API_BASE_IMPL(Real, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_EVAL_IMPL(Real, prim::kPrimReal, RealInfer, nullptr, true); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/real.h b/mindspore/core/ops/real.h index bf151f3355..af5ac148db 100644 --- a/mindspore/core/ops/real.h +++ b/mindspore/core/ops/real.h @@ -19,21 +19,18 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns a Tensor that is the real part of the input. /// Refer to Python API @ref mindspore.ops.Real for more details. -class MS_CORE_API Real : public PrimitiveC { +class MIND_API Real : public BaseOperator { public: + MIND_API_BASE_MEMBER(Real); /// \brief Constructor. - Real() : PrimitiveC(prim::kPrimReal->name()) { InitIOName({"input"}, {"output"}); } - /// \brief Destructor. - ~Real() = default; - MS_DECLARE_PARENT(Real, PrimitiveC); + Real() : BaseOperator("Real") { InitIOName({"input"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Real for the inputs. void Init() {} }; diff --git a/mindspore/core/ops/real_div.cc b/mindspore/core/ops/real_div.cc index b0981e9e2e..5cfc09d767 100644 --- a/mindspore/core/ops/real_div.cc +++ b/mindspore/core/ops/real_div.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -52,6 +53,7 @@ TypePtr RealDivInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(RealDivInferShape(primitive, input_args), RealDivInferType(primitive, input_args)); diff --git a/mindspore/core/ops/real_div.h b/mindspore/core/ops/real_div.h index ffd66fc355..93aa43efa4 100644 --- a/mindspore/core/ops/real_div.h +++ b/mindspore/core/ops/real_div.h @@ -19,28 +19,26 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameRealDiv = prim::kRealDiv; +constexpr auto kNameRealDiv = "RealDiv"; /// \brief Divides the first input tensor by the second input tensor in floating-point type element-wise. /// Refer to Python API @ref mindspore.ops.RealDiv for more details. -class MS_CORE_API RealDiv : public PrimitiveC { +class MIND_API RealDiv : public BaseOperator { public: + MIND_API_BASE_MEMBER(RealDiv); /// \brief Constructor. - RealDiv() : PrimitiveC(kNameRealDiv) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~RealDiv() = default; - MS_DECLARE_PARENT(RealDiv, PrimitiveC); + RealDiv() : BaseOperator(kNameRealDiv) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.RealDiv for the inputs. void Init() const {} }; -AbstractBasePtr RealDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RealDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reciprocal.cc b/mindspore/core/ops/reciprocal.cc index 2f4b07a5b8..b46f845f98 100644 --- a/mindspore/core/ops/reciprocal.cc +++ b/mindspore/core/ops/reciprocal.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -56,6 +57,7 @@ TypePtr ReciprocalInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(ReciprocalInferShape(primitive, input_args), diff --git a/mindspore/core/ops/reciprocal.h b/mindspore/core/ops/reciprocal.h index 3f1101d58d..4ca1f5142c 100644 --- a/mindspore/core/ops/reciprocal.h +++ b/mindspore/core/ops/reciprocal.h @@ -18,28 +18,26 @@ #define MINDSPORE_CORE_OPS_RECIPROCAL_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameReciprocal = prim::kReciprocal; +constexpr auto kNameReciprocal = "Reciprocal"; /// \brief Returns reciprocal of a tensor element-wise. /// Refer to Python API @ref mindspore.ops.Reciprocal for more details. -class MS_CORE_API Reciprocal : public PrimitiveC { +class MIND_API Reciprocal : public BaseOperator { public: + MIND_API_BASE_MEMBER(Reciprocal); /// \brief Constructor. - Reciprocal() : PrimitiveC(prim::kPrimReciprocal->name()) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Reciprocal() = default; - MS_DECLARE_PARENT(Reciprocal, PrimitiveC); + Reciprocal() : BaseOperator("Reciprocal") { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Reciprocal for the inputs. void Init() const {} }; -AbstractBasePtr ReciprocalInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReciprocalInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce.cc b/mindspore/core/ops/reduce.cc index 462b73d5b2..c32abae639 100644 --- a/mindspore/core/ops/reduce.cc +++ b/mindspore/core/ops/reduce.cc @@ -22,15 +22,17 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Reduce::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, MakeValue(keep_dims)); } +void Reduce::set_keep_dims(const bool keep_dims) { (void)this->AddAttr(kKeepDims, api::MakeValue(keep_dims)); } bool Reduce::get_keep_dims() const { return GetValue(GetAttr(kKeepDims)); } void Reduce::Init(const bool keep_dims) { this->set_keep_dims(keep_dims); } +MIND_API_BASE_IMPL(Reduce, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameReduce, Reduce); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce.h b/mindspore/core/ops/reduce.h index eaa022da1e..b5d7ce02d5 100644 --- a/mindspore/core/ops/reduce.h +++ b/mindspore/core/ops/reduce.h @@ -20,26 +20,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduce = "Reduce"; /// \brief Reduce defined Reduce operator prototype of lite. -class MS_CORE_API Reduce : public PrimitiveC { +class MIND_API Reduce : public BaseOperator { public: + MIND_API_BASE_MEMBER(Reduce); /// \brief Constructor. - Reduce() : PrimitiveC(kNameReduce) { InitIOName({"input_x", "axis"}, {"y"}); } + Reduce() : BaseOperator(kNameReduce) { InitIOName({"input_x", "axis"}, {"y"}); } /// \brief Constructor. - explicit Reduce(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"input_x", "axis"}, {"y"}); } - - /// \brief Destructor. - ~Reduce() = default; - - MS_DECLARE_PARENT(Reduce, PrimitiveC); + explicit Reduce(const std::string k_name) : BaseOperator(k_name) { InitIOName({"input_x", "axis"}, {"y"}); } /// \brief Method to init the op's attributes. /// @@ -56,8 +52,8 @@ class MS_CORE_API Reduce : public PrimitiveC { /// \return keep_dims attribute. bool get_keep_dims() const; }; -AbstractBasePtr ReduceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReduceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_all.cc b/mindspore/core/ops/reduce_all.cc index 43b252c1ec..bcec058ee9 100644 --- a/mindspore/core/ops/reduce_all.cc +++ b/mindspore/core/ops/reduce_all.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceAll, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceAll, ReduceAll); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_all.h b/mindspore/core/ops/reduce_all.h index 4ca74e4111..1b6090c7d3 100644 --- a/mindspore/core/ops/reduce_all.h +++ b/mindspore/core/ops/reduce_all.h @@ -20,22 +20,20 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceAll = "ReduceAll"; /// \brief Reduces a dimension of a tensor by the "logical AND" of all elements in the dimension. /// Refer to Python API @ref mindspore.ops.ReduceAll for more details. -class MS_CORE_API ReduceAll : public Reduce { +class MIND_API ReduceAll : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceAll); /// \brief Constructor. ReduceAll() : Reduce(kNameReduceAll) { InitIOName({"input_x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceAll() = default; - MS_DECLARE_PARENT(ReduceAll, Reduce); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_any.cc b/mindspore/core/ops/reduce_any.cc index 8f257e112e..af53adff16 100644 --- a/mindspore/core/ops/reduce_any.cc +++ b/mindspore/core/ops/reduce_any.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceAny, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceAny, ReduceAny); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_any.h b/mindspore/core/ops/reduce_any.h index 36896f411a..d81d9d6718 100644 --- a/mindspore/core/ops/reduce_any.h +++ b/mindspore/core/ops/reduce_any.h @@ -20,22 +20,20 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceAny = "ReduceAny"; /// \brief Reduces a dimension of a tensor by the "logical OR" of all elements in the dimension. /// Refer to Python API @ref mindspore.ops.ReduceAny for more details. -class MS_CORE_API ReduceAny : public Reduce { +class MIND_API ReduceAny : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceAny); /// \brief Constructor. ReduceAny() : Reduce(kNameReduceAny) { InitIOName({"input_x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceAny() = default; - MS_DECLARE_PARENT(ReduceAny, Reduce); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_asum.cc b/mindspore/core/ops/reduce_asum.cc index 028cd6f920..7775721fa6 100644 --- a/mindspore/core/ops/reduce_asum.cc +++ b/mindspore/core/ops/reduce_asum.cc @@ -18,9 +18,11 @@ #include "ops/reduce_asum.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceASum, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceASum, ReduceASum); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_asum.h b/mindspore/core/ops/reduce_asum.h index 8f568f45e6..a8e0b7aa0b 100644 --- a/mindspore/core/ops/reduce_asum.h +++ b/mindspore/core/ops/reduce_asum.h @@ -20,18 +20,17 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceASum = "ReduceASum"; -class MS_CORE_API ReduceASum : public Reduce { +class MIND_API ReduceASum : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceASum); ReduceASum() : Reduce(kNameReduceASum) { InitIOName({"input_x", "axis"}, {"y"}); } - ~ReduceASum() = default; - MS_DECLARE_PARENT(ReduceASum, Reduce); void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/reduce_max.cc b/mindspore/core/ops/reduce_max.cc index 561dbd01f8..95d163398a 100644 --- a/mindspore/core/ops/reduce_max.cc +++ b/mindspore/core/ops/reduce_max.cc @@ -18,9 +18,11 @@ #include "ops/reduce_max.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceMax, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceMax, ReduceMax); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_max.h b/mindspore/core/ops/reduce_max.h index 36f9a5d615..847341568c 100644 --- a/mindspore/core/ops/reduce_max.h +++ b/mindspore/core/ops/reduce_max.h @@ -20,22 +20,20 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceMax = "ReduceMax"; /// \brief Reduces a dimension of a tensor by the maximum value in this dimension. /// Refer to Python API @ref mindspore.ops.ReduceMax for more details. -class MS_CORE_API ReduceMax : public Reduce { +class MIND_API ReduceMax : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceMax); /// \brief Constructor. ReduceMax() : Reduce(kNameReduceMax) { InitIOName({"input_x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceMax() = default; - MS_DECLARE_PARENT(ReduceMax, Reduce); /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ReduceMax for the inputs. void Init() const {} }; diff --git a/mindspore/core/ops/reduce_mean.cc b/mindspore/core/ops/reduce_mean.cc index 267d88817e..b85aa6f3df 100644 --- a/mindspore/core/ops/reduce_mean.cc +++ b/mindspore/core/ops/reduce_mean.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceMean, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceMean, ReduceMean); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_mean.h b/mindspore/core/ops/reduce_mean.h index 59461713ef..d26c7ab720 100644 --- a/mindspore/core/ops/reduce_mean.h +++ b/mindspore/core/ops/reduce_mean.h @@ -20,22 +20,20 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceMean = "ReduceMean"; /// \brief Reduces a dimension of a tensor by averaging all elements in the dimension. /// Refer to Python API @ref mindspore.ops.ReduceMean for more details. -class MS_CORE_API ReduceMean : public Reduce { +class MIND_API ReduceMean : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceMean); /// \brief Constructor. ReduceMean() : Reduce(kNameReduceMean) { InitIOName({"input_x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceMean() = default; - MS_DECLARE_PARENT(ReduceMean, Reduce); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_min.cc b/mindspore/core/ops/reduce_min.cc index fecfc26180..023408bcba 100644 --- a/mindspore/core/ops/reduce_min.cc +++ b/mindspore/core/ops/reduce_min.cc @@ -16,9 +16,12 @@ #include "ops/reduce_min.h" #include +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceMin, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceMin, ReduceMin); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_min.h b/mindspore/core/ops/reduce_min.h index a8ddd24d94..fb92b99fbe 100644 --- a/mindspore/core/ops/reduce_min.h +++ b/mindspore/core/ops/reduce_min.h @@ -20,22 +20,20 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceMin = "ReduceMin"; /// \brief Reduces a dimension of a tensor by the minimum value in the dimension, by default. /// Refer to Python API @ref mindspore.ops.ReduceMin for more details. -class MS_CORE_API ReduceMin : public Reduce { +class MIND_API ReduceMin : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceMin); /// \brief Constructor. ReduceMin() : Reduce(kNameReduceMin) { InitIOName({"input_x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceMin() = default; - MS_DECLARE_PARENT(ReduceMin, Reduce); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_prod.cc b/mindspore/core/ops/reduce_prod.cc index 38e8b0a63b..c48981e5f1 100644 --- a/mindspore/core/ops/reduce_prod.cc +++ b/mindspore/core/ops/reduce_prod.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceProd, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceProd, ReduceProd); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_prod.h b/mindspore/core/ops/reduce_prod.h index 1db4d4670b..ad894e3287 100644 --- a/mindspore/core/ops/reduce_prod.h +++ b/mindspore/core/ops/reduce_prod.h @@ -20,22 +20,20 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceProd = "ReduceProd"; /// \brief Reduces a dimension of a tensor by multiplying all elements in the dimension, by default. /// Refer to Python API @ref mindspore.ops.ReduceProd for more details. -class MS_CORE_API ReduceProd : public Reduce { +class MIND_API ReduceProd : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceProd); /// \brief Constructor. ReduceProd() : Reduce(kNameReduceProd) { InitIOName({"input_x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceProd() = default; - MS_DECLARE_PARENT(ReduceProd, Reduce); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_scatter.cc b/mindspore/core/ops/reduce_scatter.cc index 1232fd4bb3..77ea8e619a 100644 --- a/mindspore/core/ops/reduce_scatter.cc +++ b/mindspore/core/ops/reduce_scatter.cc @@ -17,12 +17,14 @@ #include "ops/reduce_scatter.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceScatter, PrimitiveC, BaseOperator); void ReduceScatter::set_group(const string &group) { std::string g = group; - (void)this->AddAttr(kGroup, MakeValue(g)); + (void)this->AddAttr(kGroup, api::MakeValue(g)); } std::string ReduceScatter::get_group() const { auto value_ptr = GetAttr(kGroup); @@ -31,7 +33,7 @@ std::string ReduceScatter::get_group() const { void ReduceScatter::set_mode(const ReduceMode &mode) { int64_t m = mode; - (void)this->AddAttr(kMode, MakeValue(m)); + (void)this->AddAttr(kMode, api::MakeValue(m)); } ReduceMode ReduceScatter::get_mode() const { @@ -40,7 +42,7 @@ ReduceMode ReduceScatter::get_mode() const { } void ReduceScatter::set_rank_size(int rank_size) { - (void)this->AddAttr(kRankSize, MakeValue(static_cast(rank_size))); + (void)this->AddAttr(kRankSize, api::MakeValue(static_cast(rank_size))); } int ReduceScatter::get_rank_size() const { auto value_ptr = GetAttr(kRankSize); diff --git a/mindspore/core/ops/reduce_scatter.h b/mindspore/core/ops/reduce_scatter.h index 5d87b4e08e..d7d21d521a 100644 --- a/mindspore/core/ops/reduce_scatter.h +++ b/mindspore/core/ops/reduce_scatter.h @@ -20,18 +20,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceScatter = "ReduceScatter"; -class MS_CORE_API ReduceScatter : public PrimitiveC { +class MIND_API ReduceScatter : public BaseOperator { public: - ReduceScatter() : PrimitiveC(kNameReduceScatter) { InitIOName({"input_x"}, {"output"}); } - ~ReduceScatter() = default; - MS_DECLARE_PARENT(ReduceScatter, PrimitiveC); + MIND_API_BASE_MEMBER(ReduceScatter); + ReduceScatter() : BaseOperator(kNameReduceScatter) { InitIOName({"input_x"}, {"output"}); } void Init() {} void set_group(const std::string &format); std::string get_group() const; diff --git a/mindspore/core/ops/reduce_sum.cc b/mindspore/core/ops/reduce_sum.cc index 16c9828bb2..5a64c956f6 100644 --- a/mindspore/core/ops/reduce_sum.cc +++ b/mindspore/core/ops/reduce_sum.cc @@ -16,9 +16,12 @@ #include #include +#include #include "ops/reduce_sum.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -184,6 +187,7 @@ TypePtr ReduceSumInferType(const PrimitivePtr &prim, const std::vector &input_args) { const int64_t input_num = 1; diff --git a/mindspore/core/ops/reduce_sum.h b/mindspore/core/ops/reduce_sum.h index 0da53e747a..8288bdba0b 100644 --- a/mindspore/core/ops/reduce_sum.h +++ b/mindspore/core/ops/reduce_sum.h @@ -20,27 +20,25 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceSum = "ReduceSum"; /// \brief Reduces a dimension of a tensor by summing all elements in the dimension, by default. /// Refer to Python API @ref mindspore.ops.ReduceSum for more details. -class MS_CORE_API ReduceSum : public Reduce { +class MIND_API ReduceSum : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceSum); /// \brief Constructor. ReduceSum() : Reduce(kNameReduceSum) { InitIOName({"x", "axis"}, {"y"}); } - /// \brief Destructor. - ~ReduceSum() = default; - MS_DECLARE_PARENT(ReduceSum, Reduce); /// \brief Init. void Init() const {} }; -AbstractBasePtr ReduceSumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReduceSumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_sum_square.cc b/mindspore/core/ops/reduce_sum_square.cc index b28c4cbe42..fcd96bb3ce 100644 --- a/mindspore/core/ops/reduce_sum_square.cc +++ b/mindspore/core/ops/reduce_sum_square.cc @@ -18,9 +18,11 @@ #include "ops/reduce_sum_square.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReduceSumSquare, PrimitiveC, Reduce); REGISTER_PRIMITIVE_C(kNameReduceSumSquare, ReduceSumSquare); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reduce_sum_square.h b/mindspore/core/ops/reduce_sum_square.h index 49d4828987..fe5d61115a 100644 --- a/mindspore/core/ops/reduce_sum_square.h +++ b/mindspore/core/ops/reduce_sum_square.h @@ -20,18 +20,17 @@ #include #include #include + #include "ops/reduce.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReduceSumSquare = "ReduceSumSquare"; -class MS_CORE_API ReduceSumSquare : public Reduce { +class MIND_API ReduceSumSquare : public Reduce { public: + MIND_API_BASE_MEMBER(ReduceSumSquare); ReduceSumSquare() : Reduce(kNameReduceSumSquare) { InitIOName({"input_x", "axis"}, {"y"}); } - ~ReduceSumSquare() = default; - MS_DECLARE_PARENT(ReduceSumSquare, Reduce); void Init() const {} }; } // namespace ops diff --git a/mindspore/core/ops/relu.cc b/mindspore/core/ops/relu.cc index f4bcb1c210..fb1a2e9162 100644 --- a/mindspore/core/ops/relu.cc +++ b/mindspore/core/ops/relu.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -48,6 +49,8 @@ TypePtr ReLUInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = ReLUInferType(primitive, input_args); diff --git a/mindspore/core/ops/relu.h b/mindspore/core/ops/relu.h index c95df52d2a..ceb60bfdbb 100644 --- a/mindspore/core/ops/relu.h +++ b/mindspore/core/ops/relu.h @@ -19,23 +19,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameReLU = prim::kReLU; +constexpr auto kNameReLU = "ReLU"; /// \brief Computes ReLU (Rectified Linear Unit activation function) of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.ReLU for more details. -class MS_CORE_API ReLU : public PrimitiveC { +class MIND_API ReLU : public BaseOperator { public: + MIND_API_BASE_MEMBER(ReLU); /// \brief Constructor. - ReLU() : PrimitiveC(kNameReLU) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~ReLU() = default; - MS_DECLARE_PARENT(ReLU, PrimitiveC); + ReLU() : BaseOperator(kNameReLU) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/relu6.cc b/mindspore/core/ops/relu6.cc index e0f97d9e05..83d326d46d 100644 --- a/mindspore/core/ops/relu6.cc +++ b/mindspore/core/ops/relu6.cc @@ -22,6 +22,7 @@ #include "ops/relu6.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -51,6 +52,8 @@ TypePtr ReLU6InferType(const PrimitivePtr &prim, const std::vectorname()); } } // namespace + +MIND_API_BASE_IMPL(ReLU6, PrimitiveC, BaseOperator); AbstractBasePtr ReLU6Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(ReLU6InferShape(primitive, input_args), ReLU6InferType(primitive, input_args)); diff --git a/mindspore/core/ops/relu6.h b/mindspore/core/ops/relu6.h index 314de6091c..c74b287c84 100644 --- a/mindspore/core/ops/relu6.h +++ b/mindspore/core/ops/relu6.h @@ -19,28 +19,26 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameReLU6 = prim::kReLU6; +constexpr auto kNameReLU6 = "ReLU6"; /// \brief Computes ReLU (Rectified Linear Unit) upper bounded by 6 of input tensors element-wise. /// Refer to Python API @ref mindspore.ops.ReLU6 for more details. -class MS_CORE_API ReLU6 : public PrimitiveC { +class MIND_API ReLU6 : public BaseOperator { public: + MIND_API_BASE_MEMBER(ReLU6); /// \brief Constructor. - ReLU6() : PrimitiveC(kNameReLU6) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~ReLU6() = default; - MS_DECLARE_PARENT(ReLU6, PrimitiveC); + ReLU6() : BaseOperator(kNameReLU6) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr ReLU6Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReLU6Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_RELU6_H_ diff --git a/mindspore/core/ops/reluv2.cc b/mindspore/core/ops/reluv2.cc index 3a138b392c..8207fffc83 100644 --- a/mindspore/core/ops/reluv2.cc +++ b/mindspore/core/ops/reluv2.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -103,6 +104,8 @@ TypePtr ReLUV2InferType(const PrimitivePtr &prim, const std::vector(std::vector{x_type, mask_dtype}); } } // namespace + +MIND_API_BASE_IMPL(ReLUV2, PrimitiveC, BaseOperator); AbstractBasePtr ReLUV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/reluv2.h b/mindspore/core/ops/reluv2.h index 29794b7649..8926a43b52 100644 --- a/mindspore/core/ops/reluv2.h +++ b/mindspore/core/ops/reluv2.h @@ -19,30 +19,27 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameReLUV2 = prim::kReLUV2; +constexpr auto kNameReLUV2 = "ReLUV2"; /// \brief Rectified Linear Unit activation function. /// Refer to Python API @ref mindspore.ops.ReLUV2 for more details. -class MS_CORE_API ReLUV2 : public PrimitiveC { +class MIND_API ReLUV2 : public BaseOperator { public: + MIND_API_BASE_MEMBER(ReLUV2); /// \brief Constructor. - ReLUV2() : PrimitiveC(prim::kPrimReluV2->name()) { InitIOName({"x"}, {"output", "mask"}); } + ReLUV2() : BaseOperator("ReluV2") { InitIOName({"x"}, {"output", "mask"}); } /// \brief Constructor. - explicit ReLUV2(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x"}, {"output", "mask"}); } - /// \brief Destructor. - ~ReLUV2() = default; - MS_DECLARE_PARENT(ReLUV2, PrimitiveC); + explicit ReLUV2(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x"}, {"output", "mask"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr ReLUV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReLUV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reshape.cc b/mindspore/core/ops/reshape.cc index fdb0685706..dad43736c0 100644 --- a/mindspore/core/ops/reshape.cc +++ b/mindspore/core/ops/reshape.cc @@ -24,9 +24,11 @@ #include "ops/reshape.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Reshape, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameReshape, Reshape); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reshape.h b/mindspore/core/ops/reshape.h index e755047d27..7f9dac8b60 100644 --- a/mindspore/core/ops/reshape.h +++ b/mindspore/core/ops/reshape.h @@ -20,28 +20,26 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReshape = "Reshape"; /// \brief Reshapes the input tensor with the same values based on a given shape tuple. /// Refer to Python API @ref mindspore.ops.Reshape for more details. -class MS_CORE_API Reshape : public PrimitiveC { +class MIND_API Reshape : public BaseOperator { public: + MIND_API_BASE_MEMBER(Reshape); /// \brief Constructor. - Reshape() : PrimitiveC(kNameReshape) { InitIOName({"tensor", "shape"}, {"output"}); } - /// \brief Destructor. - ~Reshape() = default; - MS_DECLARE_PARENT(Reshape, PrimitiveC); + Reshape() : BaseOperator(kNameReshape) { InitIOName({"tensor", "shape"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr ReshapeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReshapeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/resize.cc b/mindspore/core/ops/resize.cc index 46fcabce21..b0219420f1 100644 --- a/mindspore/core/ops/resize.cc +++ b/mindspore/core/ops/resize.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Resize, PrimitiveC, BaseOperator); void Resize::Init(const Format format, const ResizeMethod method, const int64_t new_height, const int64_t new_width, const bool preserve_aspect_ratio, const CoordinateTransformMode coordinate_transform_mode, const float cubic_coeff, const int64_t exclude_outside, const float extrapolation_value, @@ -42,40 +44,40 @@ void Resize::Init(const Format format, const ResizeMethod method, const int64_t } void Resize::set_format(const Format format) { int64_t swi = format; - (void)this->AddAttr(kFormat, MakeValue(swi)); + (void)this->AddAttr(kFormat, api::MakeValue(swi)); } void Resize::set_method(const ResizeMethod method) { auto swi = (int64_t)method; - (void)this->AddAttr(kMethod, MakeValue(swi)); + (void)this->AddAttr(kMethod, api::MakeValue(swi)); } -void Resize::set_new_height(const int64_t new_height) { (void)this->AddAttr(kNewHeight, MakeValue(new_height)); } +void Resize::set_new_height(const int64_t new_height) { (void)this->AddAttr(kNewHeight, api::MakeValue(new_height)); } -void Resize::set_new_width(const int64_t new_width) { (void)this->AddAttr(kNewWidth, MakeValue(new_width)); } +void Resize::set_new_width(const int64_t new_width) { (void)this->AddAttr(kNewWidth, api::MakeValue(new_width)); } void Resize::set_preserve_aspect_ratio(const bool preserve_aspect_ratio) { - (void)this->AddAttr(kPreserveAspectRatio, MakeValue(preserve_aspect_ratio)); + (void)this->AddAttr(kPreserveAspectRatio, api::MakeValue(preserve_aspect_ratio)); } void Resize::set_coordinate_transform_mode(const CoordinateTransformMode coordinate_transform_mode) { int64_t swi = coordinate_transform_mode; - (void)this->AddAttr(kCoordinateTransformMode, MakeValue(swi)); + (void)this->AddAttr(kCoordinateTransformMode, api::MakeValue(swi)); } -void Resize::set_cubic_coeff(const float cubic_coeff) { (void)this->AddAttr(kCubicCoeff, MakeValue(cubic_coeff)); } +void Resize::set_cubic_coeff(const float cubic_coeff) { (void)this->AddAttr(kCubicCoeff, api::MakeValue(cubic_coeff)); } void Resize::set_exclude_outside(const int64_t exclude_outside) { - (void)this->AddAttr(kExcludeOutside, MakeValue(exclude_outside)); + (void)this->AddAttr(kExcludeOutside, api::MakeValue(exclude_outside)); } void Resize::set_extrapolation_value(const float extrapolation_value) { - (void)this->AddAttr(kExtrapolationValue, MakeValue(extrapolation_value)); + (void)this->AddAttr(kExtrapolationValue, api::MakeValue(extrapolation_value)); } void Resize::set_nearest_mode(const NearestMode nearest_mode) { int64_t swi = (int64_t)nearest_mode; - (void)this->AddAttr(kNearestMode, MakeValue(swi)); + (void)this->AddAttr(kNearestMode, api::MakeValue(swi)); } Format Resize::get_format() const { diff --git a/mindspore/core/ops/resize.h b/mindspore/core/ops/resize.h index bd117123b4..cb0bc2428c 100644 --- a/mindspore/core/ops/resize.h +++ b/mindspore/core/ops/resize.h @@ -18,23 +18,20 @@ #define MINDSPORE_CORE_OPS_RESIZE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameResize = "Resize"; /// \brief Resize defined the Resize operator prototype of lite. -class MS_CORE_API Resize : public PrimitiveC { +class MIND_API Resize : public BaseOperator { public: + MIND_API_BASE_MEMBER(Resize); /// \brief Constructor. - Resize() : PrimitiveC(kNameResize) {} - - /// \brief Destructor. - ~Resize() = default; - - MS_DECLARE_PARENT(Resize, PrimitiveC); + Resize() : BaseOperator(kNameResize) {} /// \brief Method to init the op's attributes. /// @@ -158,8 +155,8 @@ class MS_CORE_API Resize : public PrimitiveC { NearestMode get_nearest_mode() const; }; -AbstractBasePtr ResizeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ResizeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/resize_bilinear.cc b/mindspore/core/ops/resize_bilinear.cc index ed0f340bbc..f9b82866b6 100644 --- a/mindspore/core/ops/resize_bilinear.cc +++ b/mindspore/core/ops/resize_bilinear.cc @@ -22,15 +22,17 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void ResizeBilinear::set_size(const std::vector &size) { (void)this->AddAttr(kSize, MakeValue(size)); } +MIND_API_BASE_IMPL(ResizeBilinear, PrimitiveC, BaseOperator); +void ResizeBilinear::set_size(const std::vector &size) { (void)this->AddAttr(kSize, api::MakeValue(size)); } std::vector ResizeBilinear::get_size() const { return GetValue>(GetAttr(kSize)); } void ResizeBilinear::set_align_corners(const bool align_corners) { - (void)this->AddAttr(kAlignCorners, MakeValue(align_corners)); + (void)this->AddAttr(kAlignCorners, api::MakeValue(align_corners)); } bool ResizeBilinear::get_align_corners() const { diff --git a/mindspore/core/ops/resize_bilinear.h b/mindspore/core/ops/resize_bilinear.h index f06bd0ff34..70a84cb27c 100644 --- a/mindspore/core/ops/resize_bilinear.h +++ b/mindspore/core/ops/resize_bilinear.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameResizeBilinear = "ResizeBilinear"; /// \brief Resizes an image to a certain size using the bilinear interpolation. /// Refer to Python API @ref mindspore.ops.ResizeBilinear for more details. -class MS_CORE_API ResizeBilinear : public PrimitiveC { +class MIND_API ResizeBilinear : public BaseOperator { public: + MIND_API_BASE_MEMBER(ResizeBilinear); /// \brief Constructor. - ResizeBilinear() : PrimitiveC(kNameResizeBilinear) {} - /// \brief Destructor. - ~ResizeBilinear() = default; - MS_DECLARE_PARENT(ResizeBilinear, PrimitiveC); + ResizeBilinear() : BaseOperator(kNameResizeBilinear) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ResizeBilinear for the inputs. void Init(const std::vector &size, const bool align_corners = false); /// \brief Set size. @@ -51,8 +49,8 @@ class MS_CORE_API ResizeBilinear : public PrimitiveC { /// \return align_corners. bool get_align_corners() const; }; -AbstractBasePtr ResizeBilinearInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ResizeBilinearInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/resize_nearest_neighbor.cc b/mindspore/core/ops/resize_nearest_neighbor.cc index 60e8677faa..6588348c51 100644 --- a/mindspore/core/ops/resize_nearest_neighbor.cc +++ b/mindspore/core/ops/resize_nearest_neighbor.cc @@ -23,6 +23,7 @@ #include "ops/resize_nearest_neighbor.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -30,9 +31,11 @@ void ResizeNearestNeighbor::Init(const std::vector &size, const bool al this->set_size(size); this->set_align_corners(align_corners); } -void ResizeNearestNeighbor::set_size(const std::vector &size) { (void)this->AddAttr(kSize, MakeValue(size)); } +void ResizeNearestNeighbor::set_size(const std::vector &size) { + (void)this->AddAttr(kSize, api::MakeValue(size)); +} void ResizeNearestNeighbor::set_align_corners(const bool align_corners) { - (void)this->AddAttr(kAlignCorners, MakeValue(align_corners)); + (void)this->AddAttr(kAlignCorners, api::MakeValue(align_corners)); } std::vector ResizeNearestNeighbor::get_size() const { auto value_ptr = GetAttr(kSize); @@ -85,6 +88,8 @@ TypePtr ResizeNearestNeighborInferType(const PrimitivePtr &prim, const std::vect return CheckAndConvertUtils::CheckTensorTypeValid("x", input_args[0]->BuildType(), valid_types, prim->name()); } } // namespace + +MIND_API_BASE_IMPL(ResizeNearestNeighbor, PrimitiveC, BaseOperator); AbstractBasePtr ResizeNearestNeighborInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto prim_name = primitive->name(); diff --git a/mindspore/core/ops/resize_nearest_neighbor.h b/mindspore/core/ops/resize_nearest_neighbor.h index 5aa64024cd..a5c9d972d2 100644 --- a/mindspore/core/ops/resize_nearest_neighbor.h +++ b/mindspore/core/ops/resize_nearest_neighbor.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameResizeNearestNeighbor = "ResizeNearestNeighbor"; /// \brief Resizes the input tensor by using the nearest neighbor algorithm. /// Refer to Python API @ref mindspore.ops.ResizeNearestNeighbor for more details. -class MS_CORE_API ResizeNearestNeighbor : public PrimitiveC { +class MIND_API ResizeNearestNeighbor : public BaseOperator { public: + MIND_API_BASE_MEMBER(ResizeNearestNeighbor); /// \brief Constructor. - ResizeNearestNeighbor() : PrimitiveC(kNameResizeNearestNeighbor) {} - /// \brief Destructor. - ~ResizeNearestNeighbor() = default; - MS_DECLARE_PARENT(ResizeNearestNeighbor, PrimitiveC); + ResizeNearestNeighbor() : BaseOperator(kNameResizeNearestNeighbor) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ResizeNearestNeighbor for the inputs. void Init(const std::vector &size, const bool align_corners = false); /// \brief Set size. diff --git a/mindspore/core/ops/return.cc b/mindspore/core/ops/return.cc index 2a32e4d5f9..77e2f78cbb 100644 --- a/mindspore/core/ops/return.cc +++ b/mindspore/core/ops/return.cc @@ -15,9 +15,12 @@ */ #include "ops/return.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Return, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameReturn, Return); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/return.h b/mindspore/core/ops/return.h index c472d6d70f..d79b41f7db 100644 --- a/mindspore/core/ops/return.h +++ b/mindspore/core/ops/return.h @@ -16,21 +16,17 @@ #ifndef MINDSPORE_CORE_OPS_RETURN_H_ #define MINDSPORE_CORE_OPS_RETURN_H_ -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameReturn = "Return"; /// \brief Return op is the output node, which is only used in FuncGraph. -class MS_CORE_API Return : public PrimitiveC { +class MIND_API Return : public BaseOperator { public: + MIND_API_BASE_MEMBER(Return); /// \brief Constructor. - Return() : PrimitiveC(kNameReturn) {} - - /// \brief Destructor. - ~Return() = default; - - MS_DECLARE_PARENT(Return, PrimitiveC); + Return() : BaseOperator(kNameReturn) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reverse_sequence.cc b/mindspore/core/ops/reverse_sequence.cc index f449e697f9..9413c9162b 100644 --- a/mindspore/core/ops/reverse_sequence.cc +++ b/mindspore/core/ops/reverse_sequence.cc @@ -20,15 +20,19 @@ #include "ops/reverse_sequence.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ReverseSequence, PrimitiveC, BaseOperator); void ReverseSequence::Init(const int64_t seq_dim, const int64_t batch_dim) { this->set_seq_dim(seq_dim); this->set_batch_dim(batch_dim); } -void ReverseSequence::set_seq_dim(const int64_t seq_dim) { (void)this->AddAttr(kSeqDim, MakeValue(seq_dim)); } -void ReverseSequence::set_batch_dim(const int64_t batch_dim) { (void)this->AddAttr(kBatchDim, MakeValue(batch_dim)); } +void ReverseSequence::set_seq_dim(const int64_t seq_dim) { (void)this->AddAttr(kSeqDim, api::MakeValue(seq_dim)); } +void ReverseSequence::set_batch_dim(const int64_t batch_dim) { + (void)this->AddAttr(kBatchDim, api::MakeValue(batch_dim)); +} int64_t ReverseSequence::get_seq_dim() const { return GetValue(GetAttr(kSeqDim)); } int64_t ReverseSequence::get_batch_dim() const { diff --git a/mindspore/core/ops/reverse_sequence.h b/mindspore/core/ops/reverse_sequence.h index 780cc22d0c..861417c032 100644 --- a/mindspore/core/ops/reverse_sequence.h +++ b/mindspore/core/ops/reverse_sequence.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_REVERSE_SEQUENCE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReverseSequence = "ReverseSequence"; /// \brief Reverses variable length slices. /// Refer to Python API @ref mindspore.ops.ReverseSequence for more details. -class MS_CORE_API ReverseSequence : public PrimitiveC { +class MIND_API ReverseSequence : public BaseOperator { public: + MIND_API_BASE_MEMBER(ReverseSequence); /// \brief Constructor. - ReverseSequence() : PrimitiveC(kNameReverseSequence) { InitIOName({"x", "seq_lengths"}, {"y"}); } - /// \brief Destructor. - ~ReverseSequence() = default; - MS_DECLARE_PARENT(ReverseSequence, PrimitiveC); + ReverseSequence() : BaseOperator(kNameReverseSequence) { InitIOName({"x", "seq_lengths"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ReverseSequence for the inputs. void Init(const int64_t seq_dim, const int64_t batch_dim = 0); /// \brief Set seq_dim. @@ -49,8 +47,8 @@ class MS_CORE_API ReverseSequence : public PrimitiveC { /// \return batch_dim. int64_t get_batch_dim() const; }; -AbstractBasePtr ReverseSequenceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReverseSequenceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimReverseSequence = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reverse_v2.cc b/mindspore/core/ops/reverse_v2.cc index 4b5fa0b4e5..4468676c1e 100644 --- a/mindspore/core/ops/reverse_v2.cc +++ b/mindspore/core/ops/reverse_v2.cc @@ -18,16 +18,19 @@ #include "ops/op_utils.h" #include "ops/reverse_v2.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void ReverseV2::Init(const std::vector &axis) { this->set_axis(axis); } -void ReverseV2::set_axis(const std::vector &axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void ReverseV2::set_axis(const std::vector &axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } std::vector ReverseV2::get_axis() const { auto value_ptr = GetAttr(kAxis); return GetValue>(value_ptr); } +MIND_API_BASE_IMPL(ReverseV2, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameReverseV2, ReverseV2); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/reverse_v2.h b/mindspore/core/ops/reverse_v2.h index 218bbb90d0..b281daf480 100644 --- a/mindspore/core/ops/reverse_v2.h +++ b/mindspore/core/ops/reverse_v2.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameReverseV2 = "ReverseV2"; /// \brief Reverses specific dimensions of a tensor. /// Refer to Python API @ref mindspore.ops.ReverseV2 for more details. -class MS_CORE_API ReverseV2 : public PrimitiveC { +class MIND_API ReverseV2 : public BaseOperator { public: + MIND_API_BASE_MEMBER(ReverseV2); /// \brief Constructor. - ReverseV2() : PrimitiveC(kNameReverseV2) {} - /// \brief Destructor. - ~ReverseV2() = default; - MS_DECLARE_PARENT(ReverseV2, PrimitiveC); + ReverseV2() : BaseOperator(kNameReverseV2) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.ReverseV2 for the inputs. void Init(const std::vector &axis); /// \brief Set axis. @@ -46,8 +44,8 @@ class MS_CORE_API ReverseV2 : public PrimitiveC { std::vector get_axis() const; }; -AbstractBasePtr ReverseV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ReverseV2Infer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/rfft.cc b/mindspore/core/ops/rfft.cc index 9ec37bd479..71047d610e 100644 --- a/mindspore/core/ops/rfft.cc +++ b/mindspore/core/ops/rfft.cc @@ -18,15 +18,17 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void Rfft::Init(const int64_t fft_length) { this->set_fft_length(fft_length); } -void Rfft::set_fft_length(const int64_t fft_length) { (void)this->AddAttr(kFftLength, MakeValue(fft_length)); } +void Rfft::set_fft_length(const int64_t fft_length) { (void)this->AddAttr(kFftLength, api::MakeValue(fft_length)); } int64_t Rfft::get_fft_length() const { return GetValue(GetAttr(kFftLength)); } +MIND_API_BASE_IMPL(Rfft, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameRfft, Rfft); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/rfft.h b/mindspore/core/ops/rfft.h index 7cfd497b96..a82342cb2d 100644 --- a/mindspore/core/ops/rfft.h +++ b/mindspore/core/ops/rfft.h @@ -18,23 +18,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRfft = "Rfft"; /// \brief Rfft defined the operator prototype of computing discrete fourier transform of a real-valued signal. -class MS_CORE_API Rfft : public PrimitiveC { +class MIND_API Rfft : public BaseOperator { public: + MIND_API_BASE_MEMBER(Rfft); /// \brief Constructor. - Rfft() : PrimitiveC(kNameRfft) {} - - /// \brief Destructor. - ~Rfft() = default; - - MS_DECLARE_PARENT(Rfft, PrimitiveC); + Rfft() : BaseOperator(kNameRfft) {} /// \brief Method to init the op's attributes. /// @@ -51,8 +46,8 @@ class MS_CORE_API Rfft : public PrimitiveC { /// \return the FFT length. int64_t get_fft_length() const; }; -AbstractBasePtr RfftInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RfftInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/rint.cc b/mindspore/core/ops/rint.cc index 749d8ab3d9..9758f4be04 100644 --- a/mindspore/core/ops/rint.cc +++ b/mindspore/core/ops/rint.cc @@ -23,6 +23,8 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -52,6 +54,7 @@ TypePtr RintInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/rint.h b/mindspore/core/ops/rint.h index ac816cb95d..27a64da94d 100644 --- a/mindspore/core/ops/rint.h +++ b/mindspore/core/ops/rint.h @@ -19,22 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRint = "Rint"; -class Rint : public PrimitiveC { +class Rint : public BaseOperator { public: - Rint() : PrimitiveC(kNameRint) { InitIOName({"x"}, {"output"}); } - ~Rint() = default; - MS_DECLARE_PARENT(Rint, PrimitiveC); + MIND_API_BASE_MEMBER(Rint); + Rint() : BaseOperator(kNameRint) { InitIOName({"x"}, {"output"}); } }; -AbstractBasePtr RintInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RintInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimRintPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/roi_pooling.cc b/mindspore/core/ops/roi_pooling.cc index 1b0f2e522f..45ecfac4dd 100644 --- a/mindspore/core/ops/roi_pooling.cc +++ b/mindspore/core/ops/roi_pooling.cc @@ -22,21 +22,23 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void ROIPooling::set_pooled_h(const int64_t pooled_h) { (void)this->AddAttr(kPooledH, MakeValue(pooled_h)); } +MIND_API_BASE_IMPL(ROIPooling, PrimitiveC, BaseOperator); +void ROIPooling::set_pooled_h(const int64_t pooled_h) { (void)this->AddAttr(kPooledH, api::MakeValue(pooled_h)); } int64_t ROIPooling::get_pooled_h() const { return GetValue(GetAttr(kPooledH)); } -void ROIPooling::set_pooled_w(const int64_t pooled_w) { (void)this->AddAttr(kPooledW, MakeValue(pooled_w)); } +void ROIPooling::set_pooled_w(const int64_t pooled_w) { (void)this->AddAttr(kPooledW, api::MakeValue(pooled_w)); } int64_t ROIPooling::get_pooled_w() const { auto value_ptr = GetAttr(kPooledW); return GetValue(value_ptr); } -void ROIPooling::set_scale(const float scale) { (void)this->AddAttr(kScale, MakeValue(scale)); } +void ROIPooling::set_scale(const float scale) { (void)this->AddAttr(kScale, api::MakeValue(scale)); } float ROIPooling::get_scale() const { auto value_ptr = GetAttr(kScale); diff --git a/mindspore/core/ops/roi_pooling.h b/mindspore/core/ops/roi_pooling.h index d04f409732..5675dcc0d1 100644 --- a/mindspore/core/ops/roi_pooling.h +++ b/mindspore/core/ops/roi_pooling.h @@ -20,23 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameROIPooling = "ROIPooling"; /// \brief ROIPooling defined the ROIPooling operator prototype. -class MS_CORE_API ROIPooling : public PrimitiveC { +class MIND_API ROIPooling : public BaseOperator { public: + MIND_API_BASE_MEMBER(ROIPooling); /// \brief Constructor. - ROIPooling() : PrimitiveC(kNameROIPooling) {} - - /// \brief Destructor. - ~ROIPooling() = default; - - MS_DECLARE_PARENT(ROIPooling, PrimitiveC); + ROIPooling() : BaseOperator(kNameROIPooling) {} /// \brief Method to init the op's attributes. /// @@ -77,8 +73,8 @@ class MS_CORE_API ROIPooling : public PrimitiveC { /// \return the size factor. float get_scale() const; }; -AbstractBasePtr ROIPoolingInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ROIPoolingInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/roll.cc b/mindspore/core/ops/roll.cc index e3111e08f3..e08eeab3de 100644 --- a/mindspore/core/ops/roll.cc +++ b/mindspore/core/ops/roll.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -48,6 +49,7 @@ TypePtr RollInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/roll.h b/mindspore/core/ops/roll.h index e66cbcd68a..d7c525cfc0 100644 --- a/mindspore/core/ops/roll.h +++ b/mindspore/core/ops/roll.h @@ -19,25 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRoll = "Roll"; /// \brief Rolls the elements of a tensor along an axis. -class Roll : public PrimitiveC { +class Roll : public BaseOperator { public: + MIND_API_BASE_MEMBER(Roll); /// \brief Constructor. - Roll() : PrimitiveC(kNameRoll) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~Roll() = default; - MS_DECLARE_PARENT(Roll, PrimitiveC); + Roll() : BaseOperator(kNameRoll) { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr RollInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RollInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_ROLL_H_ diff --git a/mindspore/core/ops/round.cc b/mindspore/core/ops/round.cc index f249bfadba..66f0c98ef3 100644 --- a/mindspore/core/ops/round.cc +++ b/mindspore/core/ops/round.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,7 @@ TypePtr RoundInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/round.h b/mindspore/core/ops/round.h index 6ed1f739b2..f665710cec 100644 --- a/mindspore/core/ops/round.h +++ b/mindspore/core/ops/round.h @@ -19,28 +19,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRound = "Round"; /// \brief Returns half to even of a tensor element-wise. /// Refer to Python API @ref mindspore.ops.Round for more details. -class MS_CORE_API Round : public PrimitiveC { +class MIND_API Round : public BaseOperator { public: + MIND_API_BASE_MEMBER(Round); /// \brief Constructor. - Round() : PrimitiveC(kNameRound) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~Round() = default; - MS_DECLARE_PARENT(Round, PrimitiveC); + Round() : BaseOperator(kNameRound) { InitIOName({"input_x"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr RoundInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RoundInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimRoundPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/rsqrt.cc b/mindspore/core/ops/rsqrt.cc index a602f93d7b..6122820f4a 100644 --- a/mindspore/core/ops/rsqrt.cc +++ b/mindspore/core/ops/rsqrt.cc @@ -15,6 +15,12 @@ */ #include "ops/rsqrt.h" +#include +#include +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -38,6 +44,7 @@ TypePtr RsqrtInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/rsqrt.h b/mindspore/core/ops/rsqrt.h index 1699ade526..5a142a5119 100644 --- a/mindspore/core/ops/rsqrt.h +++ b/mindspore/core/ops/rsqrt.h @@ -19,27 +19,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameRsqrt = "Rsqrt"; /// \brief Computes reciprocal of square root of input tensor element-wise. /// Refer to Python API @ref mindspore.ops.Rsqrt for more details. -class MS_CORE_API Rsqrt : public PrimitiveC { +class MIND_API Rsqrt : public BaseOperator { public: + MIND_API_BASE_MEMBER(Rsqrt); /// \brief Constructor. - Rsqrt() : PrimitiveC(kNameRsqrt) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Rsqrt() = default; - MS_DECLARE_PARENT(Rsqrt, PrimitiveC); + Rsqrt() : BaseOperator(kNameRsqrt) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr RsqrtInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr RsqrtInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimRsqrtPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/scalar_summary.cc b/mindspore/core/ops/scalar_summary.cc index 055589a6e4..c09c6d4f1e 100644 --- a/mindspore/core/ops/scalar_summary.cc +++ b/mindspore/core/ops/scalar_summary.cc @@ -19,6 +19,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -34,7 +35,9 @@ abstract::ShapePtr ScalarSummaryInferShape(const PrimitivePtr &primitive, return std::make_shared(ShapeVector(1)); } } // namespace -void ScalarSummary::set_side_effect_io() { (void)this->AddAttr(kSideEffectIO, MakeValue(true)); } + +MIND_API_BASE_IMPL(ScalarSummary, PrimitiveC, BaseOperator); +void ScalarSummary::set_side_effect_io() { (void)this->AddAttr(kSideEffectIO, api::MakeValue(true)); } bool ScalarSummary::get_side_effect_io() const { auto value_ptr = GetAttr(kSideEffectIO); diff --git a/mindspore/core/ops/scalar_summary.h b/mindspore/core/ops/scalar_summary.h index 35f5ca384d..141db00f52 100644 --- a/mindspore/core/ops/scalar_summary.h +++ b/mindspore/core/ops/scalar_summary.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Outputs a scalar to a protocol buffer through a scalar summary operator. /// Refer to Python API @ref mindspore.ops.ScalarSummary for more details. -class MS_CORE_API ScalarSummary : public PrimitiveC { +class MIND_API ScalarSummary : public BaseOperator { public: + MIND_API_BASE_MEMBER(ScalarSummary); /// \brief Constructor. - ScalarSummary() : PrimitiveC(prim::kPrimScalarSummary->name()) {} - /// \brief Destructor. - ~ScalarSummary() = default; - MS_DECLARE_PARENT(ScalarSummary, PrimitiveC); + ScalarSummary() : BaseOperator("ScalarSummary") {} /// \brief Init. void Init(); /// \brief Set side_effect_io. diff --git a/mindspore/core/ops/scale.cc b/mindspore/core/ops/scale.cc index 336a496216..9e4c9810fd 100644 --- a/mindspore/core/ops/scale.cc +++ b/mindspore/core/ops/scale.cc @@ -15,12 +15,15 @@ */ #include "ops/scale.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Scale, PrimitiveC, BaseOperator); void Scale::Init(const int64_t axis) { set_axis(axis); } -void Scale::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void Scale::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } int64_t Scale::get_axis() const { return GetValue(GetAttr(kAxis)); } REGISTER_PRIMITIVE_C(kNameScale, Scale); diff --git a/mindspore/core/ops/scale.h b/mindspore/core/ops/scale.h index 6058257cf5..f0fe02f252 100644 --- a/mindspore/core/ops/scale.h +++ b/mindspore/core/ops/scale.h @@ -20,27 +20,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameScale = "Scale"; /// \brief Scale defined Scale operator prototype of lite. -class MS_CORE_API Scale : public PrimitiveC { +class MIND_API Scale : public BaseOperator { public: + MIND_API_BASE_MEMBER(Scale); /// \brief Constructor. - Scale() : PrimitiveC(kNameScale) {} + Scale() : BaseOperator(kNameScale) {} /// \brief Constructor. - explicit Scale(const std::string k_name) : PrimitiveC(k_name) {} - - /// \brief Destructor. - ~Scale() = default; - - MS_DECLARE_PARENT(Scale, PrimitiveC); + explicit Scale(const std::string k_name) : BaseOperator(k_name) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/scatter_nd.cc b/mindspore/core/ops/scatter_nd.cc index 543b48601a..6e7adca79a 100644 --- a/mindspore/core/ops/scatter_nd.cc +++ b/mindspore/core/ops/scatter_nd.cc @@ -19,9 +19,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(ScatterNd, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameScatterNd, ScatterNd); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/scatter_nd.h b/mindspore/core/ops/scatter_nd.h index 115415a6ee..aff1d5d63d 100644 --- a/mindspore/core/ops/scatter_nd.h +++ b/mindspore/core/ops/scatter_nd.h @@ -19,27 +19,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameScatterNd = "ScatterNd"; /// \brief Scatters a tensor into a new tensor depending on the specified indices. /// Refer to Python API @ref mindspore.ops.ScatterNd for more details. -class MS_CORE_API ScatterNd : public PrimitiveC { +class MIND_API ScatterNd : public BaseOperator { public: + MIND_API_BASE_MEMBER(ScatterNd); /// \brief Constructor. - ScatterNd() : PrimitiveC(kNameScatterNd) { InitIOName({"indices", "update", "shape"}, {"output"}); } - /// \brief Destructor. - ~ScatterNd() = default; - MS_DECLARE_PARENT(ScatterNd, PrimitiveC); + ScatterNd() : BaseOperator(kNameScatterNd) { InitIOName({"indices", "update", "shape"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr ScatterNdInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ScatterNdInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/scatter_nd_add.cc b/mindspore/core/ops/scatter_nd_add.cc index 59bf415c1e..8f432aad64 100644 --- a/mindspore/core/ops/scatter_nd_add.cc +++ b/mindspore/core/ops/scatter_nd_add.cc @@ -23,6 +23,7 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -71,6 +72,7 @@ TypePtr ScatterNdAddInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/scatter_nd_add.h b/mindspore/core/ops/scatter_nd_add.h index 00648c339d..a8241890bc 100644 --- a/mindspore/core/ops/scatter_nd_add.h +++ b/mindspore/core/ops/scatter_nd_add.h @@ -19,22 +19,20 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameScatterNdAdd = "ScatterNdAdd"; -class ScatterNdAdd : public PrimitiveC { +class MIND_API ScatterNdAdd : public BaseOperator { public: - ScatterNdAdd() : PrimitiveC(kNameScatterNdAdd) { InitIOName({"input_x", "indices", "updates"}, {"y"}); } - ~ScatterNdAdd() = default; - MS_DECLARE_PARENT(ScatterNdAdd, PrimitiveC); + MIND_API_BASE_MEMBER(ScatterNdAdd); + ScatterNdAdd() : BaseOperator(kNameScatterNdAdd) { InitIOName({"input_x", "indices", "updates"}, {"y"}); } }; -AbstractBasePtr ScatterNdAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ScatterNdAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimScatterNdAddPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/scatter_nd_update.cc b/mindspore/core/ops/scatter_nd_update.cc index 65096362a9..bf06c15ac7 100644 --- a/mindspore/core/ops/scatter_nd_update.cc +++ b/mindspore/core/ops/scatter_nd_update.cc @@ -23,6 +23,7 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -71,6 +72,7 @@ TypePtr ScatterNdUpdateInferType(const PrimitivePtr &primitive, const std::vecto } } // namespace +MIND_API_BASE_IMPL(ScatterNdUpdate, PrimitiveC, BaseOperator); AbstractBasePtr ScatterNdUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/scatter_nd_update.h b/mindspore/core/ops/scatter_nd_update.h index 1f7f315351..def338344a 100644 --- a/mindspore/core/ops/scatter_nd_update.h +++ b/mindspore/core/ops/scatter_nd_update.h @@ -19,27 +19,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameScatterNdUpdate = "ScatterNdUpdate"; /// \brief Updates tensor values by using input indices and value. /// Refer to Python API @ref mindspore.ops.ScatterNdUpdate for more details. -class MS_CORE_API ScatterNdUpdate : public PrimitiveC { +class MIND_API ScatterNdUpdate : public BaseOperator { public: + MIND_API_BASE_MEMBER(ScatterNdUpdate); /// \brief Constructor. - ScatterNdUpdate() : PrimitiveC(kNameScatterNdUpdate) { InitIOName({"input_x", "indices", "update"}, {"output"}); } - /// \brief Destructor. - ~ScatterNdUpdate() = default; - MS_DECLARE_PARENT(ScatterNdUpdate, PrimitiveC); + ScatterNdUpdate() : BaseOperator(kNameScatterNdUpdate) { InitIOName({"input_x", "indices", "update"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr ScatterNdUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ScatterNdUpdateInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimScatterNdUpdatePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/scatter_non_aliasing_add.cc b/mindspore/core/ops/scatter_non_aliasing_add.cc index f228f9c752..155bda8f2a 100644 --- a/mindspore/core/ops/scatter_non_aliasing_add.cc +++ b/mindspore/core/ops/scatter_non_aliasing_add.cc @@ -22,6 +22,7 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -69,6 +70,7 @@ TypePtr ScatterNonAliasingAddInferType(const PrimitivePtr &primitive, const std: } } // namespace +MIND_API_BASE_IMPL(ScatterNonAliasingAdd, PrimitiveC, BaseOperator); AbstractBasePtr ScatterNonAliasingAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/scatter_non_aliasing_add.h b/mindspore/core/ops/scatter_non_aliasing_add.h index 5fbe9ea6c7..04c6b66006 100644 --- a/mindspore/core/ops/scatter_non_aliasing_add.h +++ b/mindspore/core/ops/scatter_non_aliasing_add.h @@ -19,24 +19,22 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameScatterNonAliasingAdd = "ScatterNonAliasingAdd"; -class ScatterNonAliasingAdd : public PrimitiveC { +class MIND_API ScatterNonAliasingAdd : public BaseOperator { public: - ScatterNonAliasingAdd() : PrimitiveC(kNameScatterNonAliasingAdd) { + MIND_API_BASE_MEMBER(ScatterNonAliasingAdd); + ScatterNonAliasingAdd() : BaseOperator(kNameScatterNonAliasingAdd) { InitIOName({"input_x", "indices", "updates"}, {"y"}); } - ~ScatterNonAliasingAdd() = default; - MS_DECLARE_PARENT(ScatterNonAliasingAdd, PrimitiveC); }; -AbstractBasePtr ScatterNonAliasingAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ScatterNonAliasingAddInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimScatterNonAliasingAddPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/select.cc b/mindspore/core/ops/select.cc index e777f6b9d0..cd19a567cb 100644 --- a/mindspore/core/ops/select.cc +++ b/mindspore/core/ops/select.cc @@ -18,6 +18,8 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "utils/tensor_construct_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -194,6 +196,8 @@ ValuePtr SelectInferValue(const PrimitivePtr &prim, const std::vector #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSelect = "Select"; /// \brief Returns the selected elements, either from input x or input y, depending on the condition. /// Refer to Python API @ref mindspore.ops.Select for more details. -class MS_CORE_API Select : public PrimitiveC { +class MIND_API Select : public BaseOperator { public: + MIND_API_BASE_MEMBER(Select); /// \brief Constructor. - Select() : PrimitiveC(kNameSelect) { InitIOName({"condition", "x", "y"}, {"output"}); } - /// \brief Destructor. - ~Select() = default; - MS_DECLARE_PARENT(Select, PrimitiveC); + Select() : BaseOperator(kNameSelect) { InitIOName({"condition", "x", "y"}, {"output"}); } /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/selu.cc b/mindspore/core/ops/selu.cc index 51794d01cd..c42f0bf0db 100644 --- a/mindspore/core/ops/selu.cc +++ b/mindspore/core/ops/selu.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -51,6 +52,8 @@ TypePtr SeLUInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = SeLUInferType(primitive, input_args); diff --git a/mindspore/core/ops/selu.h b/mindspore/core/ops/selu.h index 36ae6ff72b..55d6449010 100644 --- a/mindspore/core/ops/selu.h +++ b/mindspore/core/ops/selu.h @@ -19,19 +19,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSeLU = "SeLU"; -class SeLU : public PrimitiveC { +class MIND_API SeLU : public BaseOperator { public: - SeLU() : PrimitiveC(kNameSeLU) { InitIOName({"x"}, {"output"}); } - ~SeLU() = default; - MS_DECLARE_PARENT(SeLU, PrimitiveC); + MIND_API_BASE_MEMBER(SeLU); + SeLU() : BaseOperator(kNameSeLU) { InitIOName({"x"}, {"output"}); } void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/sgd.cc b/mindspore/core/ops/sgd.cc index 50497df5fb..930904cd80 100644 --- a/mindspore/core/ops/sgd.cc +++ b/mindspore/core/ops/sgd.cc @@ -15,9 +15,13 @@ */ #include "ops/sgd.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SGD, PrimitiveC, BaseOperator); void SGD::Init(const float dampening, const float weight_decay, const bool nesterov) { set_nesterov(nesterov); set_dampening(dampening); @@ -26,12 +30,12 @@ void SGD::Init(const float dampening, const float weight_decay, const bool neste void SGD::set_dampening(const float dampening) { if (get_nesterov()) CheckAndConvertUtils::CheckValue(kDampening, dampening, kEqual, 0.0, name()); - (void)AddAttr(kDampening, MakeValue(dampening)); + (void)AddAttr(kDampening, api::MakeValue(dampening)); } -void SGD::set_weight_decay(const float weight_decay) { (void)AddAttr(kWeightDecay, MakeValue(weight_decay)); } +void SGD::set_weight_decay(const float weight_decay) { (void)AddAttr(kWeightDecay, api::MakeValue(weight_decay)); } -void SGD::set_nesterov(const bool nesterov) { (void)AddAttr(kNesterov, MakeValue(nesterov)); } +void SGD::set_nesterov(const bool nesterov) { (void)AddAttr(kNesterov, api::MakeValue(nesterov)); } float SGD::get_dampening() const { auto value_ptr = GetAttr(kDampening); diff --git a/mindspore/core/ops/sgd.h b/mindspore/core/ops/sgd.h index 91197e5434..a44ff40a64 100644 --- a/mindspore/core/ops/sgd.h +++ b/mindspore/core/ops/sgd.h @@ -18,22 +18,19 @@ #define MINDSPORE_CORE_OPS_SGD_H_ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" + +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameSGD = "SGD"; /// \brief Computes the stochastic gradient descent. /// Refer to Python API @ref mindspore.ops.SGD for more details. -class MS_CORE_API SGD : public PrimitiveC { +class MIND_API SGD : public BaseOperator { public: + MIND_API_BASE_MEMBER(SGD); /// \brief Constructor. - SGD() : PrimitiveC(kNameSGD) {} - /// \brief Destructor. - ~SGD() = default; - MS_DECLARE_PARENT(SGD, PrimitiveC); + SGD() : BaseOperator(kNameSGD) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.SGD for the inputs. void Init(const float dampening = 0.0, const float weight_decay = 0.0, const bool nesterov = false); /// \brief Set dampening. @@ -55,8 +52,8 @@ class MS_CORE_API SGD : public PrimitiveC { /// \return nesterov. bool get_nesterov() const; }; -AbstractBasePtr SGDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SGDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimSGD = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/shape.cc b/mindspore/core/ops/shape.cc index 5b98e466e3..6dfece955f 100644 --- a/mindspore/core/ops/shape.cc +++ b/mindspore/core/ops/shape.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Shape, PrimitiveC, BaseOperator); AbstractBasePtr ShapeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { // infer shape diff --git a/mindspore/core/ops/shape.h b/mindspore/core/ops/shape.h index 900ee7826a..763ca99dff 100644 --- a/mindspore/core/ops/shape.h +++ b/mindspore/core/ops/shape.h @@ -20,21 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns the shape of the input tensor. /// Refer to Python API @ref mindspore.ops.Shape for more details. -class MS_CORE_API Shape : public PrimitiveC { +class MIND_API Shape : public BaseOperator { public: + MIND_API_BASE_MEMBER(Shape); /// \brief Constructor. - Shape() : PrimitiveC(prim::kPrimShape->name()) {} - /// \brief Destructor. - ~Shape() = default; - MS_DECLARE_PARENT(Shape, PrimitiveC); + Shape() : BaseOperator("Shape") {} /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/sigmoid.cc b/mindspore/core/ops/sigmoid.cc index 5acbb55332..b5ec8799bf 100644 --- a/mindspore/core/ops/sigmoid.cc +++ b/mindspore/core/ops/sigmoid.cc @@ -18,6 +18,9 @@ #include #include #include +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +44,8 @@ TypePtr SigmoidInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/sigmoid.h b/mindspore/core/ops/sigmoid.h index ee6454e288..54462f469b 100644 --- a/mindspore/core/ops/sigmoid.h +++ b/mindspore/core/ops/sigmoid.h @@ -18,26 +18,24 @@ #define MINDSPORE_CORE_OPS_SIGMOID_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSigmoid = "Sigmoid"; /// \brief Sigmoid activation function. Refer to Python API @ref mindspore.ops.Sigmoid for more details. -class MS_CORE_API Sigmoid : public PrimitiveC { +class MIND_API Sigmoid : public BaseOperator { public: + MIND_API_BASE_MEMBER(Sigmoid); /// \brief Constructor. - Sigmoid() : PrimitiveC(kNameSigmoid) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Sigmoid() = default; - MS_DECLARE_PARENT(Sigmoid, PrimitiveC); + Sigmoid() : BaseOperator(kNameSigmoid) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr SigmoidInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SigmoidInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSigmoidPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sigmoid_cross_entropy_with_logits.cc b/mindspore/core/ops/sigmoid_cross_entropy_with_logits.cc index 74a881c9a5..0a77a163b8 100644 --- a/mindspore/core/ops/sigmoid_cross_entropy_with_logits.cc +++ b/mindspore/core/ops/sigmoid_cross_entropy_with_logits.cc @@ -17,6 +17,7 @@ #include "ops/sigmoid_cross_entropy_with_logits.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -62,6 +63,8 @@ TypePtr SigmoidCrossEntropyWithLogitsInferType(const PrimitivePtr &primitive, return logits_type; } } // namespace + +MIND_API_BASE_IMPL(SigmoidCrossEntropyWithLogits, PrimitiveC, BaseOperator); AbstractBasePtr SigmoidCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = SigmoidCrossEntropyWithLogitsInferType(primitive, input_args); diff --git a/mindspore/core/ops/sigmoid_cross_entropy_with_logits.h b/mindspore/core/ops/sigmoid_cross_entropy_with_logits.h index 5c9e4b5dd6..9920d03218 100644 --- a/mindspore/core/ops/sigmoid_cross_entropy_with_logits.h +++ b/mindspore/core/ops/sigmoid_cross_entropy_with_logits.h @@ -21,29 +21,28 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSigmoidCrossEntropyWithLogits = "SigmoidCrossEntropyWithLogits"; /// \brief Uses the given logits to compute sigmoid cross entropy between the logits and the label. /// Refer to Python API @ref mindspore.ops.SigmoidCrossEntropyWithLogits for more details. -class MS_CORE_API SigmoidCrossEntropyWithLogits : public PrimitiveC { +class MIND_API SigmoidCrossEntropyWithLogits : public BaseOperator { public: + MIND_API_BASE_MEMBER(SigmoidCrossEntropyWithLogits); /// \brief Constructor. - SigmoidCrossEntropyWithLogits() : PrimitiveC(kNameSigmoidCrossEntropyWithLogits) { + SigmoidCrossEntropyWithLogits() : BaseOperator(kNameSigmoidCrossEntropyWithLogits) { InitIOName({"predict", "target"}, {"loss"}); } - /// \brief Destructor. - ~SigmoidCrossEntropyWithLogits() = default; - MS_DECLARE_PARENT(SigmoidCrossEntropyWithLogits, PrimitiveC); /// \brief Init. void Init() const {} }; -AbstractBasePtr SigmoidCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SigmoidCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSigmoidCrossEntropyWithLogitsPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sign.cc b/mindspore/core/ops/sign.cc index be02e272a6..be4af3c079 100644 --- a/mindspore/core/ops/sign.cc +++ b/mindspore/core/ops/sign.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,7 @@ TypePtr SignInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/sign.h b/mindspore/core/ops/sign.h index 4732aaae0a..880a0e86a9 100644 --- a/mindspore/core/ops/sign.h +++ b/mindspore/core/ops/sign.h @@ -19,18 +19,15 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameSign = "Sign"; -class Sign : public PrimitiveC { +class Sign : public BaseOperator { public: - Sign() : PrimitiveC(kNameSign) { InitIOName({"x"}, {"y"}); } - ~Sign() = default; - MS_DECLARE_PARENT(Sign, PrimitiveC); + MIND_API_BASE_MEMBER(Sign); + Sign() : BaseOperator(kNameSign) { InitIOName({"x"}, {"y"}); } }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sin.cc b/mindspore/core/ops/sin.cc index 241154c3bc..874bd14916 100644 --- a/mindspore/core/ops/sin.cc +++ b/mindspore/core/ops/sin.cc @@ -18,6 +18,11 @@ #include #include #include +#include +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -40,6 +45,8 @@ TypePtr SinInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/sin.h b/mindspore/core/ops/sin.h index 03ac415a91..cca3e34630 100644 --- a/mindspore/core/ops/sin.h +++ b/mindspore/core/ops/sin.h @@ -19,27 +19,24 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSin = "Sin"; /// \brief Computes sine of the input element-wise. Refer to Python API @ref mindspore.ops.Sin for more details. -class MS_CORE_API Sin : public PrimitiveC { +class MIND_API Sin : public BaseOperator { public: + MIND_API_BASE_MEMBER(Sin); /// \brief Constructor. - Sin() : PrimitiveC(kNameSin) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Sin() = default; - MS_DECLARE_PARENT(Sin, PrimitiveC); + Sin() : BaseOperator(kNameSin) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr SinInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SinInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSinPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sinh.cc b/mindspore/core/ops/sinh.cc index b5069ea1d9..05734a601a 100644 --- a/mindspore/core/ops/sinh.cc +++ b/mindspore/core/ops/sinh.cc @@ -18,6 +18,11 @@ #include #include #include +#include +#include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +50,7 @@ TypePtr SinhInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/sinh.h b/mindspore/core/ops/sinh.h index bf9d70f8eb..9cd49c6dcf 100644 --- a/mindspore/core/ops/sinh.h +++ b/mindspore/core/ops/sinh.h @@ -18,22 +18,21 @@ #define MINDSPORE_CORE_OPS_SINH_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSinh = "Sinh"; -class Sinh : public PrimitiveC { +class Sinh : public BaseOperator { public: - Sinh() : PrimitiveC(kNameSinh) { InitIOName({"x"}, {"output"}); } - ~Sinh() = default; - MS_DECLARE_PARENT(Sinh, PrimitiveC); + MIND_API_BASE_MEMBER(Sinh); + Sinh() : BaseOperator(kNameSinh) { InitIOName({"x"}, {"output"}); } void Init() {} }; -AbstractBasePtr SinhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SinhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSinhPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/size.cc b/mindspore/core/ops/size.cc index d054f05c39..6d40ae1d43 100644 --- a/mindspore/core/ops/size.cc +++ b/mindspore/core/ops/size.cc @@ -16,9 +16,12 @@ #include #include "ops/size.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Size, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSize, Size); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/size.h b/mindspore/core/ops/size.h index 6d452bffe1..b7134e5ffc 100644 --- a/mindspore/core/ops/size.h +++ b/mindspore/core/ops/size.h @@ -18,21 +18,19 @@ #define MINDSPORE_CORE_OPS_SIZE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSize = "Size"; /// \brief Returns the size of a tensor. Refer to Python API @ref mindspore.ops.Size for more details. -class MS_CORE_API Size : public PrimitiveC { +class MIND_API Size : public BaseOperator { public: + MIND_API_BASE_MEMBER(Size); /// \brief Constructor. - Size() : PrimitiveC(kNameSize) {} - /// \brief Destructor. - ~Size() = default; - MS_DECLARE_PARENT(Size, PrimitiveC); + Size() : BaseOperator(kNameSize) {} }; } // namespace ops diff --git a/mindspore/core/ops/skip_gram.cc b/mindspore/core/ops/skip_gram.cc index 1189a9592e..fb2d853724 100644 --- a/mindspore/core/ops/skip_gram.cc +++ b/mindspore/core/ops/skip_gram.cc @@ -17,22 +17,26 @@ #include "ops/skip_gram.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void SkipGram::set_include_all_grams(const bool include_all_grams) { - (void)AddAttr(kIncludeALLGrams, MakeValue(include_all_grams)); + (void)AddAttr(kIncludeALLGrams, api::MakeValue(include_all_grams)); } bool SkipGram::get_include_all_grams() const { auto value_ptr = this->GetAttr(kIncludeALLGrams); return GetValue(value_ptr); } -void SkipGram::set_max_skip_size(const int64_t max_skip_size) { (void)AddAttr(kMaxSkipSize, MakeValue(max_skip_size)); } +void SkipGram::set_max_skip_size(const int64_t max_skip_size) { + (void)AddAttr(kMaxSkipSize, api::MakeValue(max_skip_size)); +} int64_t SkipGram::get_max_skip_size() const { auto value_ptr = this->GetAttr(kMaxSkipSize); return GetValue(value_ptr); } -void SkipGram::set_ngram_size(const int64_t ngram_size) { (void)AddAttr(kNgramSize, MakeValue(ngram_size)); } +void SkipGram::set_ngram_size(const int64_t ngram_size) { (void)AddAttr(kNgramSize, api::MakeValue(ngram_size)); } int64_t SkipGram::get_ngram_size() const { auto value_ptr = this->GetAttr(kNgramSize); return GetValue(value_ptr); @@ -43,6 +47,7 @@ void SkipGram::Init(const bool include_all_grams, const int64_t max_skip_size, c this->set_ngram_size(ngram_size); } +MIND_API_BASE_IMPL(SkipGram, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSkipGram, SkipGram); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/skip_gram.h b/mindspore/core/ops/skip_gram.h index a4efd91132..19f451f128 100644 --- a/mindspore/core/ops/skip_gram.h +++ b/mindspore/core/ops/skip_gram.h @@ -22,25 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSkipGram = "SkipGram"; /// \brief SkipGram defined the SkipGram operator prototype. -class MS_CORE_API SkipGram : public PrimitiveC { +class MIND_API SkipGram : public BaseOperator { public: + MIND_API_BASE_MEMBER(SkipGram); /// \brief Constructor. - SkipGram() : PrimitiveC(kNameSkipGram) {} - - /// \brief Destructor. - ~SkipGram() = default; - - MS_DECLARE_PARENT(SkipGram, PrimitiveC); + SkipGram() : BaseOperator(kNameSkipGram) {} /// \brief Method to init the op's attributes. /// @@ -81,8 +75,8 @@ class MS_CORE_API SkipGram : public PrimitiveC { /// \return an integer value. int64_t get_ngram_size() const; }; -AbstractBasePtr SkipGramInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SkipGramInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/slice.cc b/mindspore/core/ops/slice.cc index b6f2a45088..fe8c55108f 100644 --- a/mindspore/core/ops/slice.cc +++ b/mindspore/core/ops/slice.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -88,6 +89,7 @@ TypePtr SliceInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(SliceInferShape(primitive, input_args), SliceInferType(primitive, input_args)); diff --git a/mindspore/core/ops/slice.h b/mindspore/core/ops/slice.h index 2d45688e35..1883d9bcba 100644 --- a/mindspore/core/ops/slice.h +++ b/mindspore/core/ops/slice.h @@ -20,26 +20,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSlice = "Slice"; /// \brief Slices a tensor in the specified shape. Refer to Python API @ref mindspore.ops.Slice for more details. -class MS_CORE_API Slice : public PrimitiveC { +class MIND_API Slice : public BaseOperator { public: + MIND_API_BASE_MEMBER(Slice); /// \brief Constructor. - Slice() : PrimitiveC(kNameSlice) { InitIOName({"x", "begin", "size"}, {"output"}); } - /// \brief Destructor. - ~Slice() = default; - MS_DECLARE_PARENT(Slice, PrimitiveC); + Slice() : BaseOperator(kNameSlice) { InitIOName({"x", "begin", "size"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr SliceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SliceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimSlicePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/smooth_l1_loss.cc b/mindspore/core/ops/smooth_l1_loss.cc index 50625db04e..de0f1b2f56 100644 --- a/mindspore/core/ops/smooth_l1_loss.cc +++ b/mindspore/core/ops/smooth_l1_loss.cc @@ -22,15 +22,17 @@ #include "ops/smooth_l1_loss.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SmoothL1Loss, PrimitiveC, BaseOperator); void SmoothL1Loss::Init(const float beta) { this->set_beta(beta); } -void SmoothL1Loss::set_beta(const float beta) { (void)this->AddAttr(kBeta, MakeValue(beta)); } +void SmoothL1Loss::set_beta(const float beta) { (void)this->AddAttr(kBeta, api::MakeValue(beta)); } float SmoothL1Loss::get_beta() const { auto value_ptr = this->GetAttr(kBeta); - return GetValue(value_ptr); + return GetValue(value_ptr); } namespace { diff --git a/mindspore/core/ops/smooth_l1_loss.h b/mindspore/core/ops/smooth_l1_loss.h index d0faf93bf6..7d63a2811a 100644 --- a/mindspore/core/ops/smooth_l1_loss.h +++ b/mindspore/core/ops/smooth_l1_loss.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_SMOOTH_L1_LOSS_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSmoothL1Loss = "SmoothL1Loss"; /// \brief Computes smooth L1 loss, a robust L1 loss. /// Refer to Python API @ref mindspore.ops.SmoothL1Loss for more details. -class MS_CORE_API SmoothL1Loss : public PrimitiveC { +class MIND_API SmoothL1Loss : public BaseOperator { public: + MIND_API_BASE_MEMBER(SmoothL1Loss); /// \brief Constructor. - SmoothL1Loss() : PrimitiveC(kNameSmoothL1Loss) { InitIOName({"prediction", "target"}, {"output"}); } - /// \brief Destructor. - ~SmoothL1Loss() = default; - MS_DECLARE_PARENT(SmoothL1Loss, PrimitiveC); + SmoothL1Loss() : BaseOperator(kNameSmoothL1Loss) { InitIOName({"prediction", "target"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.SmoothL1Loss for the inputs. void Init(const float beta); /// \brief Set beta. @@ -43,8 +41,8 @@ class MS_CORE_API SmoothL1Loss : public PrimitiveC { /// \return beta. float get_beta() const; }; -AbstractBasePtr SmoothL1LossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SmoothL1LossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSmoothL1LossPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/soft_margin_loss.cc b/mindspore/core/ops/soft_margin_loss.cc index 13eb8ad1e1..77e19ceea1 100644 --- a/mindspore/core/ops/soft_margin_loss.cc +++ b/mindspore/core/ops/soft_margin_loss.cc @@ -19,6 +19,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -55,6 +56,7 @@ TypePtr SoftMarginLossInferType(const PrimitivePtr &primitive, const std::vector } } // namespace +MIND_API_BASE_IMPL(SoftMarginLoss, PrimitiveC, BaseOperator); AbstractBasePtr SoftMarginLossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { return abstract::MakeAbstract(SoftMarginLossInferShape(primitive, input_args), diff --git a/mindspore/core/ops/soft_margin_loss.h b/mindspore/core/ops/soft_margin_loss.h index 1e37fcfee7..dc9e934eee 100644 --- a/mindspore/core/ops/soft_margin_loss.h +++ b/mindspore/core/ops/soft_margin_loss.h @@ -21,26 +21,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftMarginLoss = "SoftMarginLoss"; /// \brief SoftMarginLoss operation. /// Refer to Python API @ref mindspore.ops.SoftMarginLoss for more details. -class MS_CORE_API SoftMarginLoss : public PrimitiveC { +class MIND_API SoftMarginLoss : public BaseOperator { public: + MIND_API_BASE_MEMBER(SoftMarginLoss); /// \brief Constructor. - SoftMarginLoss() : PrimitiveC(kNameSoftMarginLoss) { InitIOName({"predict", "label"}, {"loss"}); } - /// \brief Destructor. - ~SoftMarginLoss() = default; - MS_DECLARE_PARENT(SoftMarginLoss, PrimitiveC); + SoftMarginLoss() : BaseOperator(kNameSoftMarginLoss) { InitIOName({"predict", "label"}, {"loss"}); } }; -AbstractBasePtr SoftMarginLossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SoftMarginLossInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/soft_shrink.cc b/mindspore/core/ops/soft_shrink.cc index eb279dd96d..e5a6414b41 100644 --- a/mindspore/core/ops/soft_shrink.cc +++ b/mindspore/core/ops/soft_shrink.cc @@ -24,6 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -53,6 +54,7 @@ TypePtr SoftShrinkInferType(const PrimitivePtr &prim, const std::vector &input_args) { return std::make_shared(SoftShrinkInferType(primitive, input_args), diff --git a/mindspore/core/ops/soft_shrink.h b/mindspore/core/ops/soft_shrink.h index 80f7bdd826..7cba0024e3 100644 --- a/mindspore/core/ops/soft_shrink.h +++ b/mindspore/core/ops/soft_shrink.h @@ -19,26 +19,24 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftShrink = "SoftShrink"; /// \brief Applies the soft shrinkage function elementwise. /// Refer to Python API @ref mindspore.ops.SoftShrink for more details. -class MS_CORE_API SoftShrink : public PrimitiveC { +class MIND_API SoftShrink : public BaseOperator { public: + MIND_API_BASE_MEMBER(SoftShrink); /// \brief Constructor. - SoftShrink() : PrimitiveC(kNameSoftShrink) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~SoftShrink() = default; - MS_DECLARE_PARENT(SoftShrink, PrimitiveC); + SoftShrink() : BaseOperator(kNameSoftShrink) { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr SoftShrinkInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SoftShrinkInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/softmax.cc b/mindspore/core/ops/softmax.cc index 03643b3ac6..7b4d8ef885 100644 --- a/mindspore/core/ops/softmax.cc +++ b/mindspore/core/ops/softmax.cc @@ -23,10 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Softmax::set_axis(const std::vector &axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void Softmax::set_axis(const std::vector &axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } std::vector Softmax::get_axis() const { auto value_ptr = GetAttr(kAxis); @@ -78,6 +79,7 @@ TypePtr SoftMaxInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(SoftMaxInferShape(primitive, input_args), SoftMaxInferType(primitive, input_args)); diff --git a/mindspore/core/ops/softmax.h b/mindspore/core/ops/softmax.h index 21c9313250..6079ed78cc 100644 --- a/mindspore/core/ops/softmax.h +++ b/mindspore/core/ops/softmax.h @@ -21,21 +21,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftmax = "Softmax"; /// \brief Softmax operation. Refer to Python API @ref mindspore.ops.Softmax for more details. -class MS_CORE_API Softmax : public PrimitiveC { +class MIND_API Softmax : public BaseOperator { public: + MIND_API_BASE_MEMBER(Softmax); /// \brief Constructor. - Softmax() : PrimitiveC(kNameSoftmax) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Softmax() = default; - MS_DECLARE_PARENT(Softmax, PrimitiveC); + Softmax() : BaseOperator(kNameSoftmax) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Softmax for the inputs. void Init(const int64_t axis = -1); /// \brief Set axis. @@ -46,8 +44,8 @@ class MS_CORE_API Softmax : public PrimitiveC { std::vector get_axis() const; }; -AbstractBasePtr SoftmaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SoftmaxInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/softmax_cross_entropy_with_logits.cc b/mindspore/core/ops/softmax_cross_entropy_with_logits.cc index 7f1a31e200..e861a9061e 100644 --- a/mindspore/core/ops/softmax_cross_entropy_with_logits.cc +++ b/mindspore/core/ops/softmax_cross_entropy_with_logits.cc @@ -20,6 +20,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -79,6 +80,8 @@ TuplePtr SoftmaxCrossEntropyWithLogitsInferType(const PrimitivePtr &primitive, return std::make_shared(std::vector{type, type}); } } // namespace + +MIND_API_BASE_IMPL(SoftmaxCrossEntropyWithLogits, PrimitiveC, BaseOperator); AbstractBasePtr SoftmaxCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { auto infer_type = SoftmaxCrossEntropyWithLogitsInferType(primitive, input_args); diff --git a/mindspore/core/ops/softmax_cross_entropy_with_logits.h b/mindspore/core/ops/softmax_cross_entropy_with_logits.h index 8e894d7241..da06d47534 100644 --- a/mindspore/core/ops/softmax_cross_entropy_with_logits.h +++ b/mindspore/core/ops/softmax_cross_entropy_with_logits.h @@ -20,29 +20,28 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftmaxCrossEntropyWithLogits = "SoftmaxCrossEntropyWithLogits"; /// \brief Gets the softmax cross-entropy value between logits and labels with one-hot encoding. /// Refer to Python API @ref mindspore.ops.SoftmaxCrossEntropyWithLogits for more details. -class MS_CORE_API SoftmaxCrossEntropyWithLogits : public PrimitiveC { +class MIND_API SoftmaxCrossEntropyWithLogits : public BaseOperator { public: + MIND_API_BASE_MEMBER(SoftmaxCrossEntropyWithLogits); /// \brief Constructor. - SoftmaxCrossEntropyWithLogits() : PrimitiveC(kNameSoftmaxCrossEntropyWithLogits) { + SoftmaxCrossEntropyWithLogits() : BaseOperator(kNameSoftmaxCrossEntropyWithLogits) { InitIOName({"features", "labels"}, {"loss", "backprop"}); } - /// \brief Destructor. - ~SoftmaxCrossEntropyWithLogits() = default; - MS_DECLARE_PARENT(SoftmaxCrossEntropyWithLogits, PrimitiveC); /// \brief Init. void Init() const {} }; -AbstractBasePtr SoftmaxCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SoftmaxCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, + const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSoftmaxCrossEntropyWithLogitsPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/softplus.cc b/mindspore/core/ops/softplus.cc index 0f6077329a..606b4d8534 100644 --- a/mindspore/core/ops/softplus.cc +++ b/mindspore/core/ops/softplus.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -46,6 +47,7 @@ TypePtr SoftplusInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/softplus.h b/mindspore/core/ops/softplus.h index d37eb424b5..fbaf6add53 100644 --- a/mindspore/core/ops/softplus.h +++ b/mindspore/core/ops/softplus.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftplus = "Softplus"; /// \brief Softplus activation function. Refer to Python API @ref mindspore.ops.Softplus for more details. -class MS_CORE_API Softplus : public PrimitiveC { +class MIND_API Softplus : public BaseOperator { public: + MIND_API_BASE_MEMBER(Softplus); /// \brief Constructor. - Softplus() : PrimitiveC(kNameSoftplus) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Softplus() = default; - MS_DECLARE_PARENT(Softplus, PrimitiveC); + Softplus() : BaseOperator(kNameSoftplus) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/softsign.cc b/mindspore/core/ops/softsign.cc index 5c76377a4f..e66c8d3b20 100644 --- a/mindspore/core/ops/softsign.cc +++ b/mindspore/core/ops/softsign.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -49,6 +50,8 @@ TypePtr SoftsignInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto type = SoftsignInferType(primitive, input_args); diff --git a/mindspore/core/ops/softsign.h b/mindspore/core/ops/softsign.h index 36fb88d4c3..ffe8b75e45 100644 --- a/mindspore/core/ops/softsign.h +++ b/mindspore/core/ops/softsign.h @@ -19,19 +19,16 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSoftsign = "Softsign"; -class Softsign : public PrimitiveC { +class MIND_API Softsign : public BaseOperator { public: - Softsign() : PrimitiveC(kNameSoftsign) { InitIOName({"x"}, {"output"}); } - ~Softsign() = default; - MS_DECLARE_PARENT(Softsign, PrimitiveC); + MIND_API_BASE_MEMBER(Softsign); + Softsign() : BaseOperator(kNameSoftsign) { InitIOName({"x"}, {"output"}); } void Init() {} }; } // namespace ops diff --git a/mindspore/core/ops/sort.cc b/mindspore/core/ops/sort.cc index a1a609ae0e..26d378fb8a 100644 --- a/mindspore/core/ops/sort.cc +++ b/mindspore/core/ops/sort.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -52,6 +53,7 @@ TuplePtr SortInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto infertype = SortInferType(primitive, input_args); diff --git a/mindspore/core/ops/sort.h b/mindspore/core/ops/sort.h index 8c417344ce..c20d8b60e0 100644 --- a/mindspore/core/ops/sort.h +++ b/mindspore/core/ops/sort.h @@ -20,22 +20,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSort = "Sort"; -class Sort : public PrimitiveC { +class MIND_API Sort : public BaseOperator { public: - Sort() : PrimitiveC(kNameSort) { InitIOName({"x"}, {"y1", "y2"}); } - ~Sort() = default; - MS_DECLARE_PARENT(Sort, PrimitiveC); + MIND_API_BASE_MEMBER(Sort); + Sort() : BaseOperator(kNameSort) { InitIOName({"x"}, {"y1", "y2"}); } }; -AbstractBasePtr SortInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SortInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/space_to_batch.cc b/mindspore/core/ops/space_to_batch.cc index 8c5d220742..0dd532a688 100644 --- a/mindspore/core/ops/space_to_batch.cc +++ b/mindspore/core/ops/space_to_batch.cc @@ -22,11 +22,12 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void SpaceToBatch::set_paddings(const std::vector> &paddings) { - (void)this->AddAttr(kPaddings, MakeValue(paddings)); + (void)this->AddAttr(kPaddings, api::MakeValue(paddings)); int64_t h = SizeToLong(paddings.size()); int64_t w = SizeToLong(paddings[0].size()); std::vector temp_w = {2, 2}; @@ -43,7 +44,7 @@ std::vector> SpaceToBatch::get_paddings() const { return GetValue>>(value_ptr); } void SpaceToBatch::set_block_size(const std::vector block_size) { - (void)this->AddAttr(kBlockSize, MakeValue(block_size)); + (void)this->AddAttr(kBlockSize, api::MakeValue(block_size)); } std::vector SpaceToBatch::get_block_size() const { @@ -55,6 +56,7 @@ void SpaceToBatch::Init(const std::vector block_size, const std::vector this->set_block_size(block_size); } +MIND_API_BASE_IMPL(SpaceToBatch, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSpaceToBatch, SpaceToBatch); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/space_to_batch.h b/mindspore/core/ops/space_to_batch.h index b1118a1f88..cb168dad3b 100644 --- a/mindspore/core/ops/space_to_batch.h +++ b/mindspore/core/ops/space_to_batch.h @@ -21,22 +21,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSpaceToBatch = "SpaceToBatch"; /// \brief Divides spatial dimensions into blocks and combines the block size with the original batch. /// Refer to Python API @ref mindspore.ops.SpaceToBatch for more details. -class MS_CORE_API SpaceToBatch : public PrimitiveC { +class MIND_API SpaceToBatch : public BaseOperator { public: + MIND_API_BASE_MEMBER(SpaceToBatch); /// \brief Constructor. - SpaceToBatch() : PrimitiveC(kNameSpaceToBatch) {} - /// \brief Destructor. - ~SpaceToBatch() = default; - MS_DECLARE_PARENT(SpaceToBatch, PrimitiveC); + SpaceToBatch() : BaseOperator(kNameSpaceToBatch) {} /// \brief Init. Refer to the parameters of python API @ref mindspore.ops.SpaceToBatch for the inputs. void Init(const std::vector block_size, const std::vector> &paddings); /// \brief Set paddings. @@ -52,8 +50,8 @@ class MS_CORE_API SpaceToBatch : public PrimitiveC { /// \return paddings. std::vector> get_paddings() const; }; -AbstractBasePtr SpaceToBatchInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SpaceToBatchInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/space_to_batch_nd.cc b/mindspore/core/ops/space_to_batch_nd.cc index 0a3720f5ce..f9af344482 100644 --- a/mindspore/core/ops/space_to_batch_nd.cc +++ b/mindspore/core/ops/space_to_batch_nd.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -37,7 +38,7 @@ void SpaceToBatchND::set_paddings(std::vector> paddings) { (void)CheckAndConvertUtils::CheckInteger(kPaddings, paddings[i][j], kGreaterEqual, 0LL, this->name()); } } - (void)this->AddAttr(kPaddings, MakeValue(paddings)); + (void)this->AddAttr(kPaddings, api::MakeValue(paddings)); } std::vector> SpaceToBatchND::get_paddings() const { @@ -51,7 +52,7 @@ void SpaceToBatchND::set_block_shape(std::vector block_shape) { for (size_t i = 0; i < block_shape.size(); i++) { (void)CheckAndConvertUtils::CheckInteger(kBlockShape, block_shape[i], kGreaterEqual, 1LL, this->name()); } - (void)this->AddAttr(kBlockShape, MakeValue(block_shape)); + (void)this->AddAttr(kBlockShape, api::MakeValue(block_shape)); } std::vector SpaceToBatchND::get_block_shape() const { @@ -63,6 +64,7 @@ void SpaceToBatchND::Init(const std::vector block_shape, const std::vec this->set_block_shape(block_shape); } +MIND_API_BASE_IMPL(SpaceToBatchND, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSpaceToBatchND, SpaceToBatchND); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/space_to_batch_nd.h b/mindspore/core/ops/space_to_batch_nd.h index 27371d19a8..c721dc52e6 100644 --- a/mindspore/core/ops/space_to_batch_nd.h +++ b/mindspore/core/ops/space_to_batch_nd.h @@ -21,22 +21,20 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSpaceToBatchND = "SpaceToBatchND"; /// \brief Divides spatial dimensions into blocks and combines the block size with the original batch. /// Refer to Python API @ref mindspore.ops.SpaceToBatchND for more details. -class MS_CORE_API SpaceToBatchND : public PrimitiveC { +class MIND_API SpaceToBatchND : public BaseOperator { public: + MIND_API_BASE_MEMBER(SpaceToBatchND); /// \brief Constructor. - SpaceToBatchND() : PrimitiveC(kNameSpaceToBatchND) {} - /// \brief Destructor. - ~SpaceToBatchND() = default; - MS_DECLARE_PARENT(SpaceToBatchND, PrimitiveC); + SpaceToBatchND() : BaseOperator(kNameSpaceToBatchND) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.SpaceToBatchND for the inputs. void Init(const std::vector block_shape, const std::vector> paddings); /// \brief Set paddings. @@ -52,8 +50,8 @@ class MS_CORE_API SpaceToBatchND : public PrimitiveC { /// \return paddings. std::vector> get_paddings() const; }; -AbstractBasePtr SpaceToBatchNDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SpaceToBatchNDInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/space_to_depth.cc b/mindspore/core/ops/space_to_depth.cc index a20e914067..355f04f175 100644 --- a/mindspore/core/ops/space_to_depth.cc +++ b/mindspore/core/ops/space_to_depth.cc @@ -15,9 +15,13 @@ */ #include "ops/space_to_depth.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SpaceToDepth, PrimitiveC, BaseOperator); void SpaceToDepth::Init(const int64_t block_size, const Format &format) { this->set_block_size(block_size); this->set_format(format); @@ -25,7 +29,7 @@ void SpaceToDepth::Init(const int64_t block_size, const Format &format) { void SpaceToDepth::set_block_size(const int64_t block_size) { CheckAndConvertUtils::Check(kBlockSize, block_size, kGreaterEqual, 2, this->name()); - (void)AddAttr(kBlockSize, MakeValue(block_size)); + (void)AddAttr(kBlockSize, api::MakeValue(block_size)); } int64_t SpaceToDepth::get_block_size() const { @@ -35,7 +39,7 @@ int64_t SpaceToDepth::get_block_size() const { void SpaceToDepth::set_format(const Format &format) { int64_t f = format; - (void)this->AddAttr(kFormat, MakeValue(f)); + (void)this->AddAttr(kFormat, api::MakeValue(f)); } Format SpaceToDepth::get_format() const { diff --git a/mindspore/core/ops/space_to_depth.h b/mindspore/core/ops/space_to_depth.h index a6f7e230e3..57ac493fdb 100644 --- a/mindspore/core/ops/space_to_depth.h +++ b/mindspore/core/ops/space_to_depth.h @@ -19,23 +19,20 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" +#include "mindapi/base/format.h" namespace mindspore { namespace ops { constexpr auto kNameSpaceToDepth = "SpaceToDepth"; /// \brief Rearranges blocks of spatial data into depth. /// Refer to Python API @ref mindspore.ops.SpaceToDepth for more details. -class MS_CORE_API SpaceToDepth : public PrimitiveC { +class MIND_API SpaceToDepth : public BaseOperator { public: + MIND_API_BASE_MEMBER(SpaceToDepth); /// \brief Constructor. - SpaceToDepth() : PrimitiveC(kNameSpaceToDepth) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~SpaceToDepth() = default; - MS_DECLARE_PARENT(SpaceToDepth, PrimitiveC); + SpaceToDepth() : BaseOperator(kNameSpaceToDepth) { InitIOName({"x"}, {"y"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.SpaceToDepth for the inputs. void Init(const int64_t block_size, const Format &format = NCHW); /// \brief Set block_size. @@ -51,8 +48,8 @@ class MS_CORE_API SpaceToDepth : public PrimitiveC { /// \return format. Format get_format() const; }; -AbstractBasePtr SpaceToDepthInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SpaceToDepthInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_SpaceToDepth_H_ diff --git a/mindspore/core/ops/sparse_apply_adadelta.cc b/mindspore/core/ops/sparse_apply_adadelta.cc index c29dd344d7..6c306bd006 100644 --- a/mindspore/core/ops/sparse_apply_adadelta.cc +++ b/mindspore/core/ops/sparse_apply_adadelta.cc @@ -23,6 +23,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -108,6 +109,7 @@ TuplePtr SparseApplyAdadeltaInferType(const PrimitivePtr &prim, const std::vecto } } // namespace +MIND_API_BASE_IMPL(SparseApplyAdadelta, PrimitiveC, BaseOperator); AbstractBasePtr SparseApplyAdadeltaInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/sparse_apply_adadelta.h b/mindspore/core/ops/sparse_apply_adadelta.h index 9acd62165d..ccd1f46b46 100644 --- a/mindspore/core/ops/sparse_apply_adadelta.h +++ b/mindspore/core/ops/sparse_apply_adadelta.h @@ -22,24 +22,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSparseApplyAdadelta = "SparseApplyAdadelta"; -class SparseApplyAdadelta : public PrimitiveC { +class SparseApplyAdadelta : public BaseOperator { public: - SparseApplyAdadelta() : PrimitiveC(kNameSparseApplyAdadelta) { + MIND_API_BASE_MEMBER(SparseApplyAdadelta); + SparseApplyAdadelta() : BaseOperator(kNameSparseApplyAdadelta) { InitIOName({"var", "accum", "accum_updata", "lr", "rho", "grad", "indices"}, {"var", "accum", "accum_updata"}); } - ~SparseApplyAdadelta() = default; - MS_DECLARE_PARENT(SparseApplyAdadelta, PrimitiveC); }; -AbstractBasePtr SparseApplyAdadeltaInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SparseApplyAdadeltaInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimSparseApplyAdadeltaPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/sparse_apply_r_m_s_prop.cc b/mindspore/core/ops/sparse_apply_r_m_s_prop.cc index 424c0a7e7a..d26aa85f23 100644 --- a/mindspore/core/ops/sparse_apply_r_m_s_prop.cc +++ b/mindspore/core/ops/sparse_apply_r_m_s_prop.cc @@ -22,6 +22,8 @@ #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" #include "utils/tensor_construct_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -101,6 +103,7 @@ TuplePtr SparseApplyRMSPropInferType(const PrimitivePtr &prim, const std::vector } } // namespace +MIND_API_BASE_IMPL(SparseApplyRMSProp, PrimitiveC, BaseOperator); AbstractBasePtr SparseApplyRMSPropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/sparse_apply_r_m_s_prop.h b/mindspore/core/ops/sparse_apply_r_m_s_prop.h index 80f2f4e3b8..1b3bf79aca 100644 --- a/mindspore/core/ops/sparse_apply_r_m_s_prop.h +++ b/mindspore/core/ops/sparse_apply_r_m_s_prop.h @@ -22,27 +22,24 @@ #include #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSparseApplyRMSProp = "SparseApplyRMSProp"; /// \brief Update relevant entries according to the rmsprop algorithm. -class SparseApplyRMSProp : public PrimitiveC { +class SparseApplyRMSProp : public BaseOperator { public: + MIND_API_BASE_MEMBER(SparseApplyRMSProp); /// \brief Constructor. - SparseApplyRMSProp() : PrimitiveC(kNameSparseApplyRMSProp) { + SparseApplyRMSProp() : BaseOperator(kNameSparseApplyRMSProp) { InitIOName({"var", "ms", "mom", "lr", "grad", "indices"}, {"var", "ms", "mom"}); } - /// \brief Destructor. - ~SparseApplyRMSProp() = default; - MS_DECLARE_PARENT(SparseApplyRMSProp, PrimitiveC); }; -AbstractBasePtr SparseApplyRMSPropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SparseApplyRMSPropInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.cc b/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.cc index 15591f60ac..679a5deb59 100644 --- a/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.cc +++ b/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.cc @@ -22,13 +22,15 @@ #include "ops/sparse_softmax_cross_entropy_with_logits.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SparseSoftmaxCrossEntropyWithLogits, PrimitiveC, BaseOperator); void SparseSoftmaxCrossEntropyWithLogits::Init(const bool is_grad) { this->set_is_grad(is_grad); } void SparseSoftmaxCrossEntropyWithLogits::set_is_grad(const bool is_grad) { - (void)this->AddAttr(kIsGrad, MakeValue(is_grad)); + (void)this->AddAttr(kIsGrad, api::MakeValue(is_grad)); } bool SparseSoftmaxCrossEntropyWithLogits::get_is_grad() const { return GetValue(GetAttr(kIsGrad)); } diff --git a/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.h b/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.h index 5b6e2c13b5..c8a196e02f 100644 --- a/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.h +++ b/mindspore/core/ops/sparse_softmax_cross_entropy_with_logits.h @@ -18,22 +18,20 @@ #define MINDSPORE_CORE_OPS_SPARSE_SOFTMAX_CROSS_ENTROPY_WITH_LOGITS_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSparseSoftmaxCrossEntropyWithLogits = "SparseSoftmaxCrossEntropyWithLogits"; /// \brief Computes the softmax cross-entropy value between logits and sparse encoding labels. /// Refer to Python API @ref mindspore.ops.SparseSoftmaxCrossEntropyWithLogits for more details. -class MS_CORE_API SparseSoftmaxCrossEntropyWithLogits : public PrimitiveC { +class MIND_API SparseSoftmaxCrossEntropyWithLogits : public BaseOperator { public: + MIND_API_BASE_MEMBER(SparseSoftmaxCrossEntropyWithLogits); /// \brief Constructor. - SparseSoftmaxCrossEntropyWithLogits() : PrimitiveC(kNameSparseSoftmaxCrossEntropyWithLogits) {} - /// \brief Destructor. - ~SparseSoftmaxCrossEntropyWithLogits() = default; - MS_DECLARE_PARENT(SparseSoftmaxCrossEntropyWithLogits, PrimitiveC); + SparseSoftmaxCrossEntropyWithLogits() : BaseOperator(kNameSparseSoftmaxCrossEntropyWithLogits) {} /// \brief Init. /// Refer to the parameters of python API @ref mindspore.ops.SparseSoftmaxCrossEntropyWithLogits for the inputs. void Init(const bool is_grad = false); @@ -44,9 +42,9 @@ class MS_CORE_API SparseSoftmaxCrossEntropyWithLogits : public PrimitiveC { /// \return is_grad. bool get_is_grad() const; }; -AbstractBasePtr SparseSoftmaxCrossEntropyWithLogitsInfer(const abstract::AnalysisEnginePtr &, - const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SparseSoftmaxCrossEntropyWithLogitsInfer( + const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sparse_to_dense.cc b/mindspore/core/ops/sparse_to_dense.cc index 7b756bb114..9d0885cb11 100644 --- a/mindspore/core/ops/sparse_to_dense.cc +++ b/mindspore/core/ops/sparse_to_dense.cc @@ -21,9 +21,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SparseToDense, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSparseToDense, SparseToDense); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sparse_to_dense.h b/mindspore/core/ops/sparse_to_dense.h index 1181eac5c9..0f10397f94 100644 --- a/mindspore/core/ops/sparse_to_dense.h +++ b/mindspore/core/ops/sparse_to_dense.h @@ -18,27 +18,25 @@ #define MINDSPORE_CORE_OPS_SPARSE_TO_DENSE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSparseToDense = "SparseToDense"; /// \brief Converts a sparse representation into a dense tensor. /// Refer to Python API @ref mindspore.ops.SparseToDense for more details. -class MS_CORE_API SparseToDense : public PrimitiveC { +class MIND_API SparseToDense : public BaseOperator { public: + MIND_API_BASE_MEMBER(SparseToDense); /// \brief Constructor. - SparseToDense() : PrimitiveC(kNameSparseToDense) { InitIOName({"indices", "values", "dense_shape"}, {"output"}); } - /// \brief Destructor. - ~SparseToDense() = default; - MS_DECLARE_PARENT(SparseToDense, PrimitiveC); + SparseToDense() : BaseOperator(kNameSparseToDense) { InitIOName({"indices", "values", "dense_shape"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr SparseToDenseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SparseToDenseInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/splice.cc b/mindspore/core/ops/splice.cc index 1f5cda842e..4ce8148605 100644 --- a/mindspore/core/ops/splice.cc +++ b/mindspore/core/ops/splice.cc @@ -17,8 +17,11 @@ #include "ops/splice.h" #include #include "ops/op_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Splice, PrimitiveC, BaseOperator); void Splice::Init(const std::vector &contexts, const std::vector &forward_indexes, int64_t output_dims) { this->set_context(contexts); @@ -27,14 +30,14 @@ void Splice::Init(const std::vector &contexts, const std::vector &contexts) { - (void)this->AddAttr(kSpliceContext, MakeValue(contexts)); + (void)this->AddAttr(kSpliceContext, api::MakeValue(contexts)); } void Splice::set_forward_indexes(const std::vector &forward_indexes) { - (void)this->AddAttr(kSpliceForwardIndexes, MakeValue(forward_indexes)); + (void)this->AddAttr(kSpliceForwardIndexes, api::MakeValue(forward_indexes)); } -void Splice::set_output_dim(int64_t output_dim) { (void)this->AddAttr(kSpliceOutputDims, MakeValue(output_dim)); } +void Splice::set_output_dim(int64_t output_dim) { (void)this->AddAttr(kSpliceOutputDims, api::MakeValue(output_dim)); } std::vector Splice::get_context() const { auto value_ptr = GetAttr(kSpliceContext); diff --git a/mindspore/core/ops/splice.h b/mindspore/core/ops/splice.h index e94a4524f6..13928bbe05 100644 --- a/mindspore/core/ops/splice.h +++ b/mindspore/core/ops/splice.h @@ -18,22 +18,19 @@ #define MINDSPORE_CORE_OPS_SPLICE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSplice = "Splice"; /// \brief All defined All operator prototype of lite. -class MS_CORE_API Splice : public PrimitiveC { +class MIND_API Splice : public BaseOperator { public: + MIND_API_BASE_MEMBER(Splice); /// \brief Constructor. - Splice() : PrimitiveC(kNameSplice) { InitIOName({"inputs"}, {"outputs"}); } - - /// \brief Destructor. - ~Splice() = default; - MS_DECLARE_PARENT(Splice, PrimitiveC); + Splice() : BaseOperator(kNameSplice) { InitIOName({"inputs"}, {"outputs"}); } /// \brief Method to init the op's attributes. /// @@ -71,8 +68,8 @@ class MS_CORE_API Splice : public PrimitiveC { /// /// \param[in] output_dim Define the output_dim. int64_t get_output_dim() const; - AbstractBasePtr SpliceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); + abstract::AbstractBasePtr SpliceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/split.cc b/mindspore/core/ops/split.cc index 467927c595..684edc2d6c 100644 --- a/mindspore/core/ops/split.cc +++ b/mindspore/core/ops/split.cc @@ -18,19 +18,21 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Split, PrimitiveC, BaseOperator); void Split::Init(const int64_t axis, const int64_t output_num) { this->set_axis(axis); this->set_output_num(output_num); } void Split::set_size_splits(const std::vector &size_splits) { - (void)this->AddAttr(kSizeSplits, MakeValue(size_splits)); + (void)this->AddAttr(kSizeSplits, api::MakeValue(size_splits)); } -void Split::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } -void Split::set_output_num(const int64_t output_num) { (void)this->AddAttr(kOutputNum, MakeValue(output_num)); } +void Split::set_axis(const int64_t axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } +void Split::set_output_num(const int64_t output_num) { (void)this->AddAttr(kOutputNum, api::MakeValue(output_num)); } std::vector Split::get_size_splits() const { auto value_ptr = GetAttr(kSizeSplits); diff --git a/mindspore/core/ops/split.h b/mindspore/core/ops/split.h index 5487dea766..81e27850b4 100644 --- a/mindspore/core/ops/split.h +++ b/mindspore/core/ops/split.h @@ -19,22 +19,19 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSplit = "Split"; /// \brief Splits the input tensor into output_num of tensors along the given axis and output numbers. /// Refer to Python API @ref mindspore.ops.Split for more details. -class MS_CORE_API Split : public PrimitiveC { +class MIND_API Split : public BaseOperator { public: + MIND_API_BASE_MEMBER(Split); /// \brief Constructor. - Split() : PrimitiveC(kNameSplit) {} - /// \brief Destructor. - ~Split() = default; - MS_DECLARE_PARENT(Split, PrimitiveC); + Split() : BaseOperator(kNameSplit) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Split for the inputs. void Init(const int64_t axis, const int64_t output_num); /// \brief Set size_splits. @@ -56,8 +53,8 @@ class MS_CORE_API Split : public PrimitiveC { /// \return output_num. int64_t get_output_num() const; }; -AbstractBasePtr SplitInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SplitInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimSplit = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/split_v.cc b/mindspore/core/ops/split_v.cc index 0b6f30695c..b79c614631 100644 --- a/mindspore/core/ops/split_v.cc +++ b/mindspore/core/ops/split_v.cc @@ -20,6 +20,8 @@ #include "ops/op_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -96,6 +98,7 @@ TuplePtr SplitVInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/split_v.h b/mindspore/core/ops/split_v.h index df3395426b..c4e8f69337 100644 --- a/mindspore/core/ops/split_v.h +++ b/mindspore/core/ops/split_v.h @@ -19,26 +19,23 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSplitV = "SplitV"; /// \brief Splits the input tensor into num_split tensors along the given dimension. /// Refer to Python API @ref mindspore.ops.SplitV for more details. -class SplitV : public PrimitiveC { +class SplitV : public BaseOperator { public: + MIND_API_BASE_MEMBER(SplitV); /// \brief Constructor. - SplitV() : PrimitiveC(kNameSplitV) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~SplitV() = default; - MS_DECLARE_PARENT(SplitV, PrimitiveC); + SplitV() : BaseOperator(kNameSplitV) { InitIOName({"input_x"}, {"output"}); } }; -AbstractBasePtr SplitVInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SplitVInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_Split_V_H_ diff --git a/mindspore/core/ops/split_with_overlap.cc b/mindspore/core/ops/split_with_overlap.cc index d81863de74..d8de0dcc66 100644 --- a/mindspore/core/ops/split_with_overlap.cc +++ b/mindspore/core/ops/split_with_overlap.cc @@ -16,8 +16,11 @@ #include "ops/split_with_overlap.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SplitWithOverlap, PrimitiveC, BaseOperator); void SplitWithOverlap::Init(int64_t number_split, const std::vector &ratio, const std::vector &extend_top, const std::vector &extend_bottom, int64_t split_dim, int64_t stride, int64_t pad_top, bool trans_format) { @@ -31,28 +34,30 @@ void SplitWithOverlap::Init(int64_t number_split, const std::vector &ra this->set_trans_format(trans_format); } -void SplitWithOverlap::set_ratio(const std::vector &ratio) { (void)this->AddAttr(kRatio, MakeValue(ratio)); } +void SplitWithOverlap::set_ratio(const std::vector &ratio) { + (void)this->AddAttr(kRatio, api::MakeValue(ratio)); +} void SplitWithOverlap::set_extend_top(const std::vector &extend_top) { - (void)this->AddAttr(kExtendTop, MakeValue(extend_top)); + (void)this->AddAttr(kExtendTop, api::MakeValue(extend_top)); } void SplitWithOverlap::set_extend_bottom(const std::vector &extend_bottom) { - (void)this->AddAttr(kExtendBottom, MakeValue(extend_bottom)); + (void)this->AddAttr(kExtendBottom, api::MakeValue(extend_bottom)); } void SplitWithOverlap::set_number_split(int64_t number_split) { - (void)this->AddAttr(kNumberSplit, MakeValue(number_split)); + (void)this->AddAttr(kNumberSplit, api::MakeValue(number_split)); } -void SplitWithOverlap::set_split_dim(int64_t split_dim) { (void)this->AddAttr(kSplitDim, MakeValue(split_dim)); } +void SplitWithOverlap::set_split_dim(int64_t split_dim) { (void)this->AddAttr(kSplitDim, api::MakeValue(split_dim)); } -void SplitWithOverlap::set_split_stride(int64_t stride) { (void)this->AddAttr(kSplitStride, MakeValue(stride)); } +void SplitWithOverlap::set_split_stride(int64_t stride) { (void)this->AddAttr(kSplitStride, api::MakeValue(stride)); } -void SplitWithOverlap::set_pad_top(int64_t pad_top) { (void)this->AddAttr(kPadTop, MakeValue(pad_top)); } +void SplitWithOverlap::set_pad_top(int64_t pad_top) { (void)this->AddAttr(kPadTop, api::MakeValue(pad_top)); } void SplitWithOverlap::set_trans_format(bool trans_format) { - (void)this->AddAttr(kTransFormat, MakeValue(trans_format)); + (void)this->AddAttr(kTransFormat, api::MakeValue(trans_format)); } std::vector SplitWithOverlap::get_ratio() const { diff --git a/mindspore/core/ops/split_with_overlap.h b/mindspore/core/ops/split_with_overlap.h index 5b5f2bb85a..4d8a7d9f2c 100644 --- a/mindspore/core/ops/split_with_overlap.h +++ b/mindspore/core/ops/split_with_overlap.h @@ -18,20 +18,17 @@ #define MINDSPORE_CORE_OPS_SPLIT_WITH_OVERLAP_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" +#include "ops/base_operator.h" + namespace mindspore { namespace ops { constexpr auto kNameSplitWithOverlap = "SplitWithOverlap"; /// \brief All defined All operator prototype of lite. -class MS_CORE_API SplitWithOverlap : public PrimitiveC { +class MIND_API SplitWithOverlap : public BaseOperator { public: + MIND_API_BASE_MEMBER(SplitWithOverlap); /// \brief Constructor. - SplitWithOverlap() : PrimitiveC(kNameSplitWithOverlap) {} - - /// \brief Destructor. - ~SplitWithOverlap() = default; - MS_DECLARE_PARENT(SplitWithOverlap, PrimitiveC); + SplitWithOverlap() : BaseOperator(kNameSplitWithOverlap) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/sqrt.cc b/mindspore/core/ops/sqrt.cc index 609f11f6c5..52b9b475f1 100644 --- a/mindspore/core/ops/sqrt.cc +++ b/mindspore/core/ops/sqrt.cc @@ -18,9 +18,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Sqrt, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSqrt, Sqrt); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sqrt.h b/mindspore/core/ops/sqrt.h index 7b3c96f130..e4c97020fc 100644 --- a/mindspore/core/ops/sqrt.h +++ b/mindspore/core/ops/sqrt.h @@ -16,21 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_SQRT_H_ #define MINDSPORE_CORE_OPS_SQRT_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSqrt = "Sqrt"; /// \brief Returns square root of a tensor element-wise. Refer to Python API @ref mindspore.ops.Sqrt for more details. -class MS_CORE_API Sqrt : public PrimitiveC { +class MIND_API Sqrt : public BaseOperator { public: + MIND_API_BASE_MEMBER(Sqrt); /// \brief Constructor. - Sqrt() : PrimitiveC(kNameSqrt) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Sqrt() = default; - MS_DECLARE_PARENT(Sqrt, PrimitiveC); + Sqrt() : BaseOperator(kNameSqrt) { InitIOName({"x"}, {"output"}); } /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/square.cc b/mindspore/core/ops/square.cc index 58959508bf..46cc68893c 100644 --- a/mindspore/core/ops/square.cc +++ b/mindspore/core/ops/square.cc @@ -17,6 +17,8 @@ #include "ops/square.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -138,6 +140,8 @@ ValuePtr SquareInferValue(const PrimitivePtr &prim, const std::vector #include -#include "abstract/abstract_value.h" -#include "ops/primitive_c.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Returns square of a tensor element-wise. Refer to Python API @ref mindspore.ops.Square for more details. -class MS_CORE_API Square : public PrimitiveC { +class MIND_API Square : public BaseOperator { public: + MIND_API_BASE_MEMBER(Square); /// \brief Constructor. - Square() : PrimitiveC(prim::kPrimSquare->name()) { InitIOName({"input_x"}, {"output"}); } - /// \brief Destructor. - ~Square() = default; - MS_DECLARE_PARENT(Square, PrimitiveC); + Square() : BaseOperator("Square") { InitIOName({"input_x"}, {"output"}); } /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/squared_difference.cc b/mindspore/core/ops/squared_difference.cc index 16bae5881c..bee3a99e92 100644 --- a/mindspore/core/ops/squared_difference.cc +++ b/mindspore/core/ops/squared_difference.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -65,6 +66,7 @@ TypePtr SquaredDifferenceInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/squared_difference.h b/mindspore/core/ops/squared_difference.h index fc46020084..946f80f3c8 100644 --- a/mindspore/core/ops/squared_difference.h +++ b/mindspore/core/ops/squared_difference.h @@ -18,27 +18,25 @@ #define MINDSPORE_CORE_OPS_SQUARED_DIFFERENCE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSquaredDifference = "SquaredDifference"; /// \brief Subtracts the second input tensor from the first input tensor element-wise and returns square of it. /// Refer to Python API @ref mindspore.ops.SquaredDifference for more details. -class MS_CORE_API SquaredDifference : public PrimitiveC { +class MIND_API SquaredDifference : public BaseOperator { public: + MIND_API_BASE_MEMBER(SquaredDifference); /// \brief Constructor. - SquaredDifference() : PrimitiveC(kNameSquaredDifference) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~SquaredDifference() = default; - MS_DECLARE_PARENT(SquaredDifference, PrimitiveC); + SquaredDifference() : BaseOperator(kNameSquaredDifference) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr SquaredDifferenceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SquaredDifferenceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimSquaredDifferencePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/squeeze.cc b/mindspore/core/ops/squeeze.cc index 84510b29fb..0c1a5181fc 100644 --- a/mindspore/core/ops/squeeze.cc +++ b/mindspore/core/ops/squeeze.cc @@ -15,11 +15,15 @@ */ #include "ops/squeeze.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { void Squeeze::Init(const std::vector &axis) { set_axis(axis); } -void Squeeze::set_axis(const std::vector &axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void Squeeze::set_axis(const std::vector &axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } std::vector Squeeze::get_axis() const { return GetValue>(GetAttr(kAxis)); } namespace { @@ -81,6 +85,8 @@ TypePtr SqueezeInferType(const PrimitivePtr &prim, const std::vectorBuildType(); } } // namespace + +MIND_API_BASE_IMPL(Squeeze, PrimitiveC, BaseOperator); AbstractBasePtr SqueezeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { const size_t x_index = 0; diff --git a/mindspore/core/ops/squeeze.h b/mindspore/core/ops/squeeze.h index bdf9e3877f..f0249df153 100644 --- a/mindspore/core/ops/squeeze.h +++ b/mindspore/core/ops/squeeze.h @@ -22,24 +22,20 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSqueeze = "Squeeze"; /// \brief Returns a tensor with the same data type but dimensions of 1 are removed based on axis. /// Refer to Python API @ref mindspore.ops.Squeeze for more details. -class MS_CORE_API Squeeze : public PrimitiveC { +class MIND_API Squeeze : public BaseOperator { public: + MIND_API_BASE_MEMBER(Squeeze); /// \brief Constructor. - Squeeze() : PrimitiveC(kNameSqueeze) { InitIOName({"x"}, {"output"}); } - /// \brief Destructor. - ~Squeeze() = default; - MS_DECLARE_PARENT(Squeeze, PrimitiveC); + Squeeze() : BaseOperator(kNameSqueeze) { InitIOName({"x"}, {"output"}); } /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Squeeze for the inputs. void Init(const std::vector &axis = {}); /// \brief Set axis. @@ -50,8 +46,8 @@ class MS_CORE_API Squeeze : public PrimitiveC { std::vector get_axis() const; }; -AbstractBasePtr SqueezeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SqueezeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/stack.cc b/mindspore/core/ops/stack.cc index 594943456a..00711efdbd 100644 --- a/mindspore/core/ops/stack.cc +++ b/mindspore/core/ops/stack.cc @@ -15,15 +15,20 @@ */ #include "ops/stack.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void Stack::set_axis(const int64_t axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void Stack::set_axis(const int64_t axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } int64_t Stack::get_axis() const { return GetValue(GetAttr(kAxis)); } void Stack::Init(const int64_t axis) { this->set_axis(axis); } +MIND_API_BASE_IMPL(Stack, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameStack, Stack); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/stack.h b/mindspore/core/ops/stack.h index ba5f7b9f86..ac029fca9f 100644 --- a/mindspore/core/ops/stack.h +++ b/mindspore/core/ops/stack.h @@ -22,23 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameStack = "Stack"; /// \brief Stacks a list of tensors in specified axis. Refer to Python API @ref mindspore.ops.Tile for more details. -class MS_CORE_API Stack : public PrimitiveC { +class MIND_API Stack : public BaseOperator { public: + MIND_API_BASE_MEMBER(Stack); /// \brief Constructor. - Stack() : PrimitiveC(kNameStack) {} - /// \brief Destructor. - ~Stack() = default; - MS_DECLARE_PARENT(Stack, PrimitiveC); + Stack() : BaseOperator(kNameStack) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Stack for the inputs. void Init(const int64_t axis); /// \brief Set axis. @@ -48,8 +44,8 @@ class MS_CORE_API Stack : public PrimitiveC { /// \return axis. int64_t get_axis() const; }; -AbstractBasePtr StackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr StackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore #endif // MINDSPORE_CORE_OPS_STACK_H_ diff --git a/mindspore/core/ops/strided_slice.cc b/mindspore/core/ops/strided_slice.cc index bddaa8f032..fb32aa7de1 100644 --- a/mindspore/core/ops/strided_slice.cc +++ b/mindspore/core/ops/strided_slice.cc @@ -24,6 +24,8 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -392,9 +394,10 @@ TypePtr StridedSliceInferType(const PrimitivePtr &primitive, const std::vectorname()); - (void)this->AddAttr(kBeginMask, MakeValue(begin_mask)); + (void)this->AddAttr(kBeginMask, api::MakeValue(begin_mask)); } int64_t StridedSlice::get_begin_mask() const { auto value_ptr = GetAttr(kBeginMask); @@ -402,7 +405,7 @@ int64_t StridedSlice::get_begin_mask() const { } void StridedSlice::set_end_mask(int64_t end_mask) { (void)CheckAndConvertUtils::CheckInteger(kEndMask, end_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kEndMask, MakeValue(end_mask)); + (void)this->AddAttr(kEndMask, api::MakeValue(end_mask)); } int64_t StridedSlice::get_end_mask() const { auto value_ptr = GetAttr(kEndMask); @@ -416,7 +419,7 @@ void StridedSlice::set_ellipsis_mask(int64_t ellipsis_mask) { buffer << "For" << this->name() << ", only support one ellipsis in the index, but got " << this->get_end_mask(); MS_EXCEPTION(ValueError) << buffer.str(); } - (void)this->AddAttr(kEllipsisMask, MakeValue(ellipsis_mask)); + (void)this->AddAttr(kEllipsisMask, api::MakeValue(ellipsis_mask)); } int64_t StridedSlice::get_ellipsis_mask() const { auto value_ptr = GetAttr(kEllipsisMask); @@ -424,7 +427,7 @@ int64_t StridedSlice::get_ellipsis_mask() const { } void StridedSlice::set_new_axis_mask(int64_t new_axis_mask) { (void)CheckAndConvertUtils::CheckInteger(kNewAxisMask, new_axis_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kNewAxisMask, MakeValue(new_axis_mask)); + (void)this->AddAttr(kNewAxisMask, api::MakeValue(new_axis_mask)); } int64_t StridedSlice::get_new_axis_mask() const { auto value_ptr = GetAttr(kNewAxisMask); @@ -432,7 +435,7 @@ int64_t StridedSlice::get_new_axis_mask() const { } void StridedSlice::set_shrink_axis_mask(int64_t shrink_axis_mask) { (void)CheckAndConvertUtils::CheckInteger(kShrinkAxisMask, shrink_axis_mask, kGreaterEqual, 0, this->name()); - (void)this->AddAttr(kShrinkAxisMask, MakeValue(shrink_axis_mask)); + (void)this->AddAttr(kShrinkAxisMask, api::MakeValue(shrink_axis_mask)); } int64_t StridedSlice::get_shrink_axis_mask() const { auto value_ptr = GetAttr(kShrinkAxisMask); diff --git a/mindspore/core/ops/strided_slice.h b/mindspore/core/ops/strided_slice.h index 6717fe6f6d..0dfde09cc9 100644 --- a/mindspore/core/ops/strided_slice.h +++ b/mindspore/core/ops/strided_slice.h @@ -20,23 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameStridedSlice = prim::kStridedSlice; +constexpr auto kNameStridedSlice = "StridedSlice"; /// \brief Extracts a strided slice of a tensor. Refer to Python API @ref mindspore.ops.StridedSlice for more details. -class MS_CORE_API StridedSlice : public PrimitiveC { +class MIND_API StridedSlice : public BaseOperator { public: + MIND_API_BASE_MEMBER(StridedSlice); /// \brief Constructor. - StridedSlice() : PrimitiveC(prim::kPrimStridedSlice->name()) { - InitIOName({"x", "begin", "end", "strides"}, {"output"}); - } - /// \brief Destructor. - ~StridedSlice() = default; - MS_DECLARE_PARENT(StridedSlice, PrimitiveC); + StridedSlice() : BaseOperator("StridedSlice") { InitIOName({"x", "begin", "end", "strides"}, {"output"}); } /// \brief Init. Refer to the parameters of python API @ref mindspore.ops.StridedSlice for the inputs. void Init(int64_t begin_mask = 0, int64_t end_mask = 0, int64_t ellipsis_mask = 0, int64_t new_axis_mask = 0, int64_t shrink_axis_mask = 0); @@ -72,8 +68,8 @@ class MS_CORE_API StridedSlice : public PrimitiveC { int64_t get_shrink_axis_mask() const; }; -AbstractBasePtr StridedSliceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr StridedSliceInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimStridedSlicePtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/sub.cc b/mindspore/core/ops/sub.cc index 12dab47aa5..66d90ba7eb 100644 --- a/mindspore/core/ops/sub.cc +++ b/mindspore/core/ops/sub.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -52,6 +53,7 @@ TypePtr SubInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto shape = SubInferShape(primitive, input_args); diff --git a/mindspore/core/ops/sub.h b/mindspore/core/ops/sub.h index fc865f65a5..12c0930d86 100644 --- a/mindspore/core/ops/sub.h +++ b/mindspore/core/ops/sub.h @@ -20,29 +20,26 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameSub = prim::kSub; +constexpr auto kNameSub = "Sub"; /// \brief Subtracts the second input tensor from the first input tensor element-wise. /// Refer to Python API @ref mindspore.ops.Sub for more details. -class MS_CORE_API Sub : public PrimitiveC { +class MIND_API Sub : public BaseOperator { public: + MIND_API_BASE_MEMBER(Sub); /// \brief Constructor. - Sub() : PrimitiveC(kNameSub) { InitIOName({"x", "y"}, {"output"}); } - explicit Sub(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Sub() = default; - MS_DECLARE_PARENT(Sub, PrimitiveC); + Sub() : BaseOperator(kNameSub) { InitIOName({"x", "y"}, {"output"}); } + explicit Sub(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr SubInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr SubInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/switch.cc b/mindspore/core/ops/switch.cc index 8537728148..4f18710534 100644 --- a/mindspore/core/ops/switch.cc +++ b/mindspore/core/ops/switch.cc @@ -18,9 +18,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Switch, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSwitch, Switch); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/switch.h b/mindspore/core/ops/switch.h index 8ce925fc24..d534b1e86b 100644 --- a/mindspore/core/ops/switch.h +++ b/mindspore/core/ops/switch.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_SWITCH_H_ #define MINDSPORE_CORE_OPS_SWITCH_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSwitch = "Switch"; /// \brief Switch defined Switch operator prototype of lite. -class MS_CORE_API Switch : public PrimitiveC { +class MIND_API Switch : public BaseOperator { public: + MIND_API_BASE_MEMBER(Switch); /// \brief Constructor. - Switch() : PrimitiveC(kNameSwitch) {} - - /// \brief Destructor. - ~Switch() = default; - - MS_DECLARE_PARENT(Switch, PrimitiveC); + Switch() : BaseOperator(kNameSwitch) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/switch_layer.cc b/mindspore/core/ops/switch_layer.cc index 98f7aa811d..b85f974886 100644 --- a/mindspore/core/ops/switch_layer.cc +++ b/mindspore/core/ops/switch_layer.cc @@ -18,9 +18,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(SwitchLayer, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameSwitchLayer, SwitchLayer); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/switch_layer.h b/mindspore/core/ops/switch_layer.h index 50ef74e62e..ea4e67ee26 100644 --- a/mindspore/core/ops/switch_layer.h +++ b/mindspore/core/ops/switch_layer.h @@ -16,23 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_SWITCH_LAYER_H_ #define MINDSPORE_CORE_OPS_SWITCH_LAYER_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameSwitchLayer = "switch_layer"; /// \brief SwitchLayer defined SwitchLayer operator prototype of lite. -class MS_CORE_API SwitchLayer : public PrimitiveC { +class MIND_API SwitchLayer : public BaseOperator { public: + MIND_API_BASE_MEMBER(SwitchLayer); /// \brief Constructor. - SwitchLayer() : PrimitiveC(kNameSwitchLayer) {} - - /// \brief Destructor. - ~SwitchLayer() = default; - - MS_DECLARE_PARENT(SwitchLayer, PrimitiveC); + SwitchLayer() : BaseOperator(kNameSwitchLayer) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tan.cc b/mindspore/core/ops/tan.cc index ccfb8a73b7..2418d10dc5 100644 --- a/mindspore/core/ops/tan.cc +++ b/mindspore/core/ops/tan.cc @@ -22,6 +22,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -45,6 +46,8 @@ TypePtr TanInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/tan.h b/mindspore/core/ops/tan.h index 9622e46bfd..67b884ced9 100644 --- a/mindspore/core/ops/tan.h +++ b/mindspore/core/ops/tan.h @@ -18,26 +18,24 @@ #define MINDSPORE_CORE_OPS_TAN_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTan = "Tan"; /// \brief Computes tangent of x element-wise. Refer to Python API @ref mindspore.ops.Tan for more details. -class MS_CORE_API Tan : public PrimitiveC { +class MIND_API Tan : public BaseOperator { public: + MIND_API_BASE_MEMBER(Tan); /// \brief Constructor. - Tan() : PrimitiveC(kNameTan) {} - /// \brief Destructor. - ~Tan() = default; - MS_DECLARE_PARENT(Tan, PrimitiveC); + Tan() : BaseOperator(kNameTan) {} /// \brief Init. void Init() const {} }; -AbstractBasePtr TanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TanInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using kPrimTanPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tanh.cc b/mindspore/core/ops/tanh.cc index d89d88857b..cd85ca7ea6 100644 --- a/mindspore/core/ops/tanh.cc +++ b/mindspore/core/ops/tanh.cc @@ -24,6 +24,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -49,6 +50,8 @@ TypePtr TanhInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/tanh.h b/mindspore/core/ops/tanh.h index ef45b019b4..56e35e3368 100644 --- a/mindspore/core/ops/tanh.h +++ b/mindspore/core/ops/tanh.h @@ -21,27 +21,23 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTanh = "Tanh"; /// \brief Tanh activation function. Refer to Python API @ref mindspore.ops.Tanh for more details. -class MS_CORE_API Tanh : public PrimitiveC { +class MIND_API Tanh : public BaseOperator { public: + MIND_API_BASE_MEMBER(Tanh); /// \brief Constructor. - Tanh() : PrimitiveC(kNameTanh) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~Tanh() = default; - MS_DECLARE_PARENT(Tanh, PrimitiveC); + Tanh() : BaseOperator(kNameTanh) { InitIOName({"x"}, {"y"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr TanhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TanhInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimTanhPtr = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_array.cc b/mindspore/core/ops/tensor_array.cc index ccb8dfd075..00489d5937 100644 --- a/mindspore/core/ops/tensor_array.cc +++ b/mindspore/core/ops/tensor_array.cc @@ -17,9 +17,11 @@ #include "ops/tensor_array.h" #include #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorArray, PrimitiveC, BaseOperator); constexpr auto kTensorArrayDynamicSize = "dynamic_size"; constexpr auto kTensorArrayIdenticalElementShapes = "identical_element_shapes"; constexpr auto kTensorArrayElementShape = "element_shape"; @@ -34,18 +36,18 @@ void TensorArray::Init(bool dynamic_size, bool identical_element_shapes, const s } void TensorArray::set_dynamic_size(bool dynamic_size) { - (void)this->AddAttr(kTensorArrayDynamicSize, MakeValue(dynamic_size)); + (void)this->AddAttr(kTensorArrayDynamicSize, api::MakeValue(dynamic_size)); } void TensorArray::set_identical_element_shapes(bool identical_element_shapes) { - (void)this->AddAttr(kTensorArrayIdenticalElementShapes, MakeValue(identical_element_shapes)); + (void)this->AddAttr(kTensorArrayIdenticalElementShapes, api::MakeValue(identical_element_shapes)); } void TensorArray::set_element_shape(const std::vector &element_shape) { - (void)this->AddAttr(kTensorArrayElementShape, MakeValue(element_shape)); + (void)this->AddAttr(kTensorArrayElementShape, api::MakeValue(element_shape)); } -void TensorArray::set_data_type(int data_type) { (void)this->AddAttr(kTensorArrayDataType, MakeValue(data_type)); } +void TensorArray::set_data_type(int data_type) { (void)this->AddAttr(kTensorArrayDataType, api::MakeValue(data_type)); } bool TensorArray::get_dynamic_size() const { auto value_ptr = GetAttr(kTensorArrayDynamicSize); @@ -59,12 +61,15 @@ bool TensorArray::get_identical_element_shapes() const { const std::vector TensorArray::get_element_shape() const { auto value_ptr = GetAttr(kTensorArrayElementShape); - return GetValue>(value_ptr); + auto tmp = GetValue>(value_ptr); + std::vector res(tmp.begin(), tmp.end()); + return res; } int TensorArray::get_data_type() const { auto value_ptr = GetAttr(kTensorArrayDataType); - return GetValue(value_ptr); + auto tmp = GetValue(value_ptr); + return static_cast(tmp); } REGISTER_PRIMITIVE_C(kNameTensorArray, TensorArray); diff --git a/mindspore/core/ops/tensor_array.h b/mindspore/core/ops/tensor_array.h index e07d7ba8bc..c5a257e1c0 100644 --- a/mindspore/core/ops/tensor_array.h +++ b/mindspore/core/ops/tensor_array.h @@ -18,21 +18,18 @@ #define MINDSPORE_CORE_OPS_TENSOR_ARRAY_H_ #include #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { - constexpr auto kNameTensorArray = "TensorArray"; /// \brief Assert defined TensorArray operator prototype of lite. -class MS_CORE_API TensorArray : public PrimitiveC { +class MIND_API TensorArray : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorArray); /// \brief Constructor. - TensorArray() : PrimitiveC(kNameTensorArray) { InitIOName({"size"}, {"handle", "flow"}); } - /// \brief Destructor. - ~TensorArray() = default; - MS_DECLARE_PARENT(TensorArray, PrimitiveC); + TensorArray() : BaseOperator(kNameTensorArray) { InitIOName({"size"}, {"handle", "flow"}); } /// \brief Method to init the op's attributes. void Init(bool dynamic_size, bool identical_element_shapes, const std::vector &element_shape, int data_type); /// \brief Method to set dynamic_size attributes. diff --git a/mindspore/core/ops/tensor_array_read.cc b/mindspore/core/ops/tensor_array_read.cc index 83f41cb17e..f99049a650 100644 --- a/mindspore/core/ops/tensor_array_read.cc +++ b/mindspore/core/ops/tensor_array_read.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorArrayRead, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameTensorArrayRead, TensorArrayRead); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_array_read.h b/mindspore/core/ops/tensor_array_read.h index e573416153..927ade09fe 100644 --- a/mindspore/core/ops/tensor_array_read.h +++ b/mindspore/core/ops/tensor_array_read.h @@ -18,21 +18,18 @@ #define MINDSPORE_CORE_OPS_TENSOR_ARRAY_READ_H_ #include #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { - constexpr auto kNameTensorArrayRead = "TensorArrayRead"; /// \brief Assert defined TensorArrayRead operator prototype of lite. -class MS_CORE_API TensorArrayRead : public PrimitiveC { +class MIND_API TensorArrayRead : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorArrayRead); /// \brief Constructor. - TensorArrayRead() : PrimitiveC(kNameTensorArrayRead) { InitIOName({"handle", "index", "flow_in"}, {"tensor"}); } - /// \brief Destructor. - ~TensorArrayRead() = default; - MS_DECLARE_PARENT(TensorArrayRead, PrimitiveC); + TensorArrayRead() : BaseOperator(kNameTensorArrayRead) { InitIOName({"handle", "index", "flow_in"}, {"tensor"}); } /// \brief Method to init the op's attributes. void Init() const {} }; diff --git a/mindspore/core/ops/tensor_array_write.cc b/mindspore/core/ops/tensor_array_write.cc index dfdef04914..2802d91562 100644 --- a/mindspore/core/ops/tensor_array_write.cc +++ b/mindspore/core/ops/tensor_array_write.cc @@ -23,9 +23,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorArrayWrite, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameTensorArrayWrite, TensorArrayWrite); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_array_write.h b/mindspore/core/ops/tensor_array_write.h index 8221a01524..1929af9688 100644 --- a/mindspore/core/ops/tensor_array_write.h +++ b/mindspore/core/ops/tensor_array_write.h @@ -18,23 +18,20 @@ #define MINDSPORE_CORE_OPS_TENSOR_ARRAY_WRITE_H_ #include #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { - constexpr auto kNameTensorArrayWrite = "TensorArrayWrite"; /// \brief Assert defined TensorArrayWrite operator prototype of lite. -class MS_CORE_API TensorArrayWrite : public PrimitiveC { +class MIND_API TensorArrayWrite : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorArrayWrite); /// \brief Constructor. - TensorArrayWrite() : PrimitiveC(kNameTensorArrayWrite) { + TensorArrayWrite() : BaseOperator(kNameTensorArrayWrite) { InitIOName({"handle", "index", "value", "flow_in"}, {"flow_out"}); } - /// \brief Destructor. - ~TensorArrayWrite() = default; - MS_DECLARE_PARENT(TensorArrayWrite, PrimitiveC); /// \brief Method to init the op's attributes. void Init() const {} }; diff --git a/mindspore/core/ops/tensor_list_from_tensor.cc b/mindspore/core/ops/tensor_list_from_tensor.cc index f8c4d972d6..3ff4ed1cd9 100644 --- a/mindspore/core/ops/tensor_list_from_tensor.cc +++ b/mindspore/core/ops/tensor_list_from_tensor.cc @@ -17,6 +17,7 @@ #include "ops/tensor_list_from_tensor.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -36,13 +37,14 @@ int64_t TensorListFromTensor::get_shape_type() const { } void TensorListFromTensor::set_element_dtype(const int64_t element_dtype) { - (void)this->AddAttr(kElement_dtype, MakeValue(element_dtype)); + (void)this->AddAttr(kElement_dtype, api::MakeValue(element_dtype)); } void TensorListFromTensor::set_shape_type(const int64_t shape_type) { - (void)this->AddAttr(kShapeType, MakeValue(shape_type)); + (void)this->AddAttr(kShapeType, api::MakeValue(shape_type)); } +MIND_API_BASE_IMPL(TensorListFromTensor, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameTensorListFromTensor, TensorListFromTensor); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_list_from_tensor.h b/mindspore/core/ops/tensor_list_from_tensor.h index f83fc973d9..77f81e5e90 100644 --- a/mindspore/core/ops/tensor_list_from_tensor.h +++ b/mindspore/core/ops/tensor_list_from_tensor.h @@ -18,23 +18,19 @@ #define MINDSPORE_CORE_OPS_TENSOR_LIST_FROM_TENSOR_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTensorListFromTensor = "TensorListFromTensor"; /// \brief TensorListFromTensor defined TensorListFromTensor operator prototype of lite. -class MS_CORE_API TensorListFromTensor : public PrimitiveC { +class MIND_API TensorListFromTensor : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorListFromTensor); /// \brief Constructor. - TensorListFromTensor() : PrimitiveC(kNameTensorListFromTensor) {} - - /// \brief Destructor. - ~TensorListFromTensor() = default; - - MS_DECLARE_PARENT(TensorListFromTensor, PrimitiveC); + TensorListFromTensor() : BaseOperator(kNameTensorListFromTensor) {} /// \brief Method to init the op's attributes. /// @@ -58,8 +54,8 @@ class MS_CORE_API TensorListFromTensor : public PrimitiveC { /// \brief Method to get the op's shape_type attributes. int64_t get_shape_type() const; }; -AbstractBasePtr TensorListFromTensorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TensorListFromTensorInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_list_get_item.cc b/mindspore/core/ops/tensor_list_get_item.cc index 2353ff7627..9e7f5c8aad 100644 --- a/mindspore/core/ops/tensor_list_get_item.cc +++ b/mindspore/core/ops/tensor_list_get_item.cc @@ -17,13 +17,15 @@ #include "ops/tensor_list_get_item.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorListGetItem, PrimitiveC, BaseOperator); void TensorListGetItem::Init(const int64_t element_dtype) { this->set_element_dtype(element_dtype); } void TensorListGetItem::set_element_dtype(const int64_t element_dtype) { - (void)this->AddAttr(kElement_dtype, MakeValue(element_dtype)); + (void)this->AddAttr(kElement_dtype, api::MakeValue(element_dtype)); } int64_t TensorListGetItem::get_element_dtype() const { diff --git a/mindspore/core/ops/tensor_list_get_item.h b/mindspore/core/ops/tensor_list_get_item.h index 905cb7d16d..5f2ee5ec0b 100644 --- a/mindspore/core/ops/tensor_list_get_item.h +++ b/mindspore/core/ops/tensor_list_get_item.h @@ -17,22 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_TENSOR_LIST_GET_ITEM_H_ #define MINDSPORE_CORE_OPS_TENSOR_LIST_GET_ITEM_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTensorListGetItem = "TensorListGetItem"; /// \brief TensorListGetItem defined TensorListGetItem operator prototype of lite. -class MS_CORE_API TensorListGetItem : public PrimitiveC { +class MIND_API TensorListGetItem : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorListGetItem); /// \brief Constructor. - TensorListGetItem() : PrimitiveC(kNameTensorListGetItem) {} - /// \brief Destructor. - ~TensorListGetItem() = default; - - MS_DECLARE_PARENT(TensorListGetItem, PrimitiveC); + TensorListGetItem() : BaseOperator(kNameTensorListGetItem) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/tensor_list_reserve.cc b/mindspore/core/ops/tensor_list_reserve.cc index a5fc1a5f71..4d5c8c77c6 100644 --- a/mindspore/core/ops/tensor_list_reserve.cc +++ b/mindspore/core/ops/tensor_list_reserve.cc @@ -17,20 +17,22 @@ #include "ops/tensor_list_reserve.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorListReserve, PrimitiveC, BaseOperator); void TensorListReserve::Init(const int64_t element_dtype, const int64_t shape_type) { this->set_element_dtype(element_dtype); this->set_shape_type(shape_type); } void TensorListReserve::set_element_dtype(const int64_t element_dtype) { - (void)this->AddAttr(kElement_dtype, MakeValue(element_dtype)); + (void)this->AddAttr(kElement_dtype, api::MakeValue(element_dtype)); } void TensorListReserve::set_shape_type(const int64_t shape_type) { - (void)this->AddAttr(kShapeType, MakeValue(shape_type)); + (void)this->AddAttr(kShapeType, api::MakeValue(shape_type)); } int64_t TensorListReserve::get_element_dtype() const { diff --git a/mindspore/core/ops/tensor_list_reserve.h b/mindspore/core/ops/tensor_list_reserve.h index f43c70eca6..4c5b235d85 100644 --- a/mindspore/core/ops/tensor_list_reserve.h +++ b/mindspore/core/ops/tensor_list_reserve.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_TENSOR_LIST_RESERVE_H_ #define MINDSPORE_CORE_OPS_TENSOR_LIST_RESERVE_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTensorListReserve = "TensorListReserve"; /// \brief TensorListReserve defined TensorListReserve operator prototype of lite. -class MS_CORE_API TensorListReserve : public PrimitiveC { +class MIND_API TensorListReserve : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorListReserve); /// \brief Constructor. - TensorListReserve() : PrimitiveC(kNameTensorListReserve) {} - - /// \brief Destructor. - ~TensorListReserve() = default; - - MS_DECLARE_PARENT(TensorListReserve, PrimitiveC); + TensorListReserve() : BaseOperator(kNameTensorListReserve) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/tensor_list_set_item.cc b/mindspore/core/ops/tensor_list_set_item.cc index d1653ad8f6..eb23a08f04 100644 --- a/mindspore/core/ops/tensor_list_set_item.cc +++ b/mindspore/core/ops/tensor_list_set_item.cc @@ -17,13 +17,15 @@ #include "ops/tensor_list_set_item.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorListSetItem, PrimitiveC, BaseOperator); void TensorListSetItem::Init(const int64_t element_dtype) { this->set_element_dtype(element_dtype); } void TensorListSetItem::set_element_dtype(const int64_t element_dtype) { - (void)this->AddAttr(kElement_dtype, MakeValue(element_dtype)); + (void)this->AddAttr(kElement_dtype, api::MakeValue(element_dtype)); } int64_t TensorListSetItem::get_element_dtype() const { diff --git a/mindspore/core/ops/tensor_list_set_item.h b/mindspore/core/ops/tensor_list_set_item.h index f960bc69bd..a4eae4a9c9 100644 --- a/mindspore/core/ops/tensor_list_set_item.h +++ b/mindspore/core/ops/tensor_list_set_item.h @@ -17,23 +17,18 @@ #ifndef MINDSPORE_CORE_OPS_TENSOR_LIST_SET_ITEM_H_ #define MINDSPORE_CORE_OPS_TENSOR_LIST_SET_ITEM_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTensorListSetItem = "TensorListSetItem"; /// \brief TensorListSetItem defined TensorListSetItem operator prototype of lite. -class MS_CORE_API TensorListSetItem : public PrimitiveC { +class MIND_API TensorListSetItem : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorListSetItem); /// \brief Constructor. - TensorListSetItem() : PrimitiveC(kNameTensorListSetItem) {} - - /// \brief Destructor. - ~TensorListSetItem() = default; - - MS_DECLARE_PARENT(TensorListSetItem, PrimitiveC); + TensorListSetItem() : BaseOperator(kNameTensorListSetItem) {} /// \brief Method to init the op's attributes. /// diff --git a/mindspore/core/ops/tensor_list_stack.cc b/mindspore/core/ops/tensor_list_stack.cc index 00edac4595..66fcbd9001 100644 --- a/mindspore/core/ops/tensor_list_stack.cc +++ b/mindspore/core/ops/tensor_list_stack.cc @@ -20,20 +20,22 @@ #include "ops/tensor_list_stack.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorListStack, PrimitiveC, BaseOperator); void TensorListStack::Init(const int64_t num_elements, const int64_t element_dtype) { this->set_num_elements(num_elements); this->set_element_dtype(element_dtype); } void TensorListStack::set_num_elements(const int64_t num_elements) { - (void)this->AddAttr(kNumElements, MakeValue(num_elements)); + (void)this->AddAttr(kNumElements, api::MakeValue(num_elements)); } void TensorListStack::set_element_dtype(const int64_t element_dtype) { - (void)this->AddAttr(kElement_dtype, MakeValue(element_dtype)); + (void)this->AddAttr(kElement_dtype, api::MakeValue(element_dtype)); } int64_t TensorListStack::get_num_elements() const { diff --git a/mindspore/core/ops/tensor_list_stack.h b/mindspore/core/ops/tensor_list_stack.h index 7b18b48a94..86be78dbe0 100644 --- a/mindspore/core/ops/tensor_list_stack.h +++ b/mindspore/core/ops/tensor_list_stack.h @@ -18,23 +18,19 @@ #define MINDSPORE_CORE_OPS_TENSOR_LIST_STACK_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTensorListStack = "TensorListStack"; /// \brief TensorListStack defined TensorListStack operator prototype of lite. -class MS_CORE_API TensorListStack : public PrimitiveC { +class MIND_API TensorListStack : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorListStack); /// \brief Constructor. - TensorListStack() : PrimitiveC(kNameTensorListStack) {} - - /// \brief Destructor. - ~TensorListStack() = default; - - MS_DECLARE_PARENT(TensorListStack, PrimitiveC); + TensorListStack() : BaseOperator(kNameTensorListStack) {} /// \brief Method to init the op's attributes. /// @@ -59,8 +55,8 @@ class MS_CORE_API TensorListStack : public PrimitiveC { int64_t get_element_dtype() const; }; -AbstractBasePtr TensorListStackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TensorListStackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_shape.cc b/mindspore/core/ops/tensor_shape.cc index 19eaa25fdc..5a05aa0b90 100644 --- a/mindspore/core/ops/tensor_shape.cc +++ b/mindspore/core/ops/tensor_shape.cc @@ -18,8 +18,12 @@ #include #include "ops/dynamic_shape.h" #include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TensorShape, PrimitiveC, BaseOperator); abstract::AbstractBasePtr TensorShapeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, const std::vector &input_args) { CheckAndConvertUtils::CheckInputArgs(input_args, kEqual, 1, primitive->name()); diff --git a/mindspore/core/ops/tensor_shape.h b/mindspore/core/ops/tensor_shape.h index 248e30ed3a..84a7712ae8 100644 --- a/mindspore/core/ops/tensor_shape.h +++ b/mindspore/core/ops/tensor_shape.h @@ -15,15 +15,16 @@ */ #ifndef MINDSPORE_CORE_OPS_TENSOR_SHAPE_H_ #define MINDSPORE_CORE_OPS_TENSOR_SHAPE_H_ -#include "ops/primitive_c.h" -#include "base/core_ops.h" + +#include "ops/base_operator.h" + namespace mindspore { namespace ops { -class TensorShape : public PrimitiveC { +constexpr auto kNameTensorShape = "TensorShape"; +class MIND_API TensorShape : public BaseOperator { public: - TensorShape() : PrimitiveC(prim::kPrimTensorShape->name()) {} - ~TensorShape() = default; - MS_DECLARE_PARENT(TensorShape, PrimitiveC); + MIND_API_BASE_MEMBER(TensorShape); + TensorShape() : BaseOperator(kNameTensorShape) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tensor_summary.cc b/mindspore/core/ops/tensor_summary.cc index 884e1a3c08..86c188c006 100644 --- a/mindspore/core/ops/tensor_summary.cc +++ b/mindspore/core/ops/tensor_summary.cc @@ -19,6 +19,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -34,7 +35,9 @@ abstract::ShapePtr TensorSummaryInferShape(const PrimitivePtr &primitive, return std::make_shared(ShapeVector(1)); } } // namespace -void TensorSummary::set_side_effect_io() { (void)this->AddAttr(kSideEffectIO, MakeValue(true)); } + +MIND_API_BASE_IMPL(TensorSummary, PrimitiveC, BaseOperator); +void TensorSummary::set_side_effect_io() { (void)this->AddAttr(kSideEffectIO, api::MakeValue(true)); } bool TensorSummary::get_side_effect_io() const { auto value_ptr = GetAttr(kSideEffectIO); diff --git a/mindspore/core/ops/tensor_summary.h b/mindspore/core/ops/tensor_summary.h index 3add733dc2..5516b7e500 100644 --- a/mindspore/core/ops/tensor_summary.h +++ b/mindspore/core/ops/tensor_summary.h @@ -20,22 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Outputs a tensor to a protocol buffer through a tensor summary operator. /// Refer to Python API @ref mindspore.ops.TensorSummary for more details. -class MS_CORE_API TensorSummary : public PrimitiveC { +class MIND_API TensorSummary : public BaseOperator { public: + MIND_API_BASE_MEMBER(TensorSummary); /// \brief Constructor. - TensorSummary() : PrimitiveC(prim::kPrimTensorSummary->name()) {} - /// \brief Destructor. - ~TensorSummary() = default; - MS_DECLARE_PARENT(TensorSummary, PrimitiveC); + TensorSummary() : BaseOperator("TensorSummary") {} /// \brief Init. void Init(); /// \brief Set side_effect_io. diff --git a/mindspore/core/ops/tile.cc b/mindspore/core/ops/tile.cc index 440fd80a67..e81883ff31 100644 --- a/mindspore/core/ops/tile.cc +++ b/mindspore/core/ops/tile.cc @@ -19,6 +19,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -105,6 +106,7 @@ TypePtr TileInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(TileInferShape(primitive, input_args), TileInferType(primitive, input_args)); diff --git a/mindspore/core/ops/tile.h b/mindspore/core/ops/tile.h index 19acb43b18..162dc07d4f 100644 --- a/mindspore/core/ops/tile.h +++ b/mindspore/core/ops/tile.h @@ -21,27 +21,24 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameTile = prim::kTile; +constexpr auto kNameTile = "Tile"; /// \brief Replicates a tensor with given multiples times. Refer to Python API @ref mindspore.ops.Tile for more details. -class MS_CORE_API Tile : public PrimitiveC { +class MIND_API Tile : public BaseOperator { public: + MIND_API_BASE_MEMBER(Tile); /// \brief Constructor. - Tile() : PrimitiveC(kNameTile) { InitIOName({"x", "multiples"}, {"output"}); } - explicit Tile(const std::string k_name) : PrimitiveC(k_name) { InitIOName({"x", "multiples"}, {"output"}); } - /// \brief Destructor. - ~Tile() = default; - MS_DECLARE_PARENT(Tile, PrimitiveC); + Tile() : BaseOperator(kNameTile) { InitIOName({"x", "multiples"}, {"output"}); } + explicit Tile(const std::string k_name) : BaseOperator(k_name) { InitIOName({"x", "multiples"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr TileInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TileInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/to_format.cc b/mindspore/core/ops/to_format.cc index d4501b0701..688c6fa099 100644 --- a/mindspore/core/ops/to_format.cc +++ b/mindspore/core/ops/to_format.cc @@ -23,16 +23,18 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { -void ToFormat::set_src_t(const int64_t src_t) { (void)this->AddAttr(kSrcT, MakeValue(src_t)); } +MIND_API_BASE_IMPL(ToFormat, PrimitiveC, BaseOperator); +void ToFormat::set_src_t(const int64_t src_t) { (void)this->AddAttr(kSrcT, api::MakeValue(src_t)); } int64_t ToFormat::get_src_t() const { auto value_ptr = GetAttr(kSrcT); return GetValue(value_ptr); } -void ToFormat::set_dst_t(const int64_t dst_t) { (void)this->AddAttr(kDstT, MakeValue(dst_t)); } +void ToFormat::set_dst_t(const int64_t dst_t) { (void)this->AddAttr(kDstT, api::MakeValue(dst_t)); } int64_t ToFormat::get_dst_t() const { auto value_ptr = GetAttr(kDstT); return GetValue(value_ptr); diff --git a/mindspore/core/ops/to_format.h b/mindspore/core/ops/to_format.h index 141d285ceb..0606f9f6cb 100644 --- a/mindspore/core/ops/to_format.h +++ b/mindspore/core/ops/to_format.h @@ -20,18 +20,17 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameToFormat = "ToFormat"; -class MS_CORE_API ToFormat : public PrimitiveC { +class MIND_API ToFormat : public BaseOperator { public: - ToFormat() : PrimitiveC(kNameToFormat) {} - ~ToFormat() = default; - MS_DECLARE_PARENT(ToFormat, PrimitiveC); + MIND_API_BASE_MEMBER(ToFormat); + ToFormat() : BaseOperator(kNameToFormat) {} void Init(const int64_t src_t, const int64_t dst_t); void set_src_t(const int64_t src_t); void set_dst_t(const int64_t dst_t); diff --git a/mindspore/core/ops/topk.cc b/mindspore/core/ops/topk.cc index 03db601089..dac5e82519 100644 --- a/mindspore/core/ops/topk.cc +++ b/mindspore/core/ops/topk.cc @@ -15,14 +15,17 @@ */ #include +#include #include "ops/topk.h" #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TopK, PrimitiveC, BaseOperator); void TopK::Init(const bool sorted) { this->set_sorted(sorted); } -void TopK::set_sorted(const bool sorted) { (void)this->AddAttr(kSorted, MakeValue(sorted)); } +void TopK::set_sorted(const bool sorted) { (void)this->AddAttr(kSorted, api::MakeValue(sorted)); } bool TopK::get_sorted() const { auto value_ptr = this->GetAttr(kSorted); diff --git a/mindspore/core/ops/topk.h b/mindspore/core/ops/topk.h index 6ce734c260..e966226cfe 100644 --- a/mindspore/core/ops/topk.h +++ b/mindspore/core/ops/topk.h @@ -19,24 +19,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTopK = "TopK"; /// \brief Finds values and indices of the k largest entries along the last dimension. /// Refer to Python API @ref mindspore.ops.TopK for more details. -class MS_CORE_API TopK : public PrimitiveC { +class MIND_API TopK : public BaseOperator { public: + MIND_API_BASE_MEMBER(TopK); /// \brief Constructor. - explicit TopK(const std::string &k_name = kNameTopK) : PrimitiveC(k_name) { + explicit TopK(const std::string &k_name = kNameTopK) : BaseOperator(k_name) { InitIOName({"input", "k"}, {"values", "indices"}); } - /// \brief Destructor. - ~TopK() = default; - MS_DECLARE_PARENT(TopK, PrimitiveC); /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.TopK for the inputs. void Init(const bool sorted = false); /// \brief Set sorted. @@ -46,8 +44,8 @@ class MS_CORE_API TopK : public PrimitiveC { /// \return sorted. bool get_sorted() const; }; -AbstractBasePtr TopKInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TopKInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/transpose.cc b/mindspore/core/ops/transpose.cc index d5b1b08d29..e54f4ce70c 100644 --- a/mindspore/core/ops/transpose.cc +++ b/mindspore/core/ops/transpose.cc @@ -21,6 +21,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -94,6 +95,7 @@ TypePtr TransposeInferType(const PrimitivePtr &prim, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/transpose.h b/mindspore/core/ops/transpose.h index 7060d547e7..03d6e32439 100644 --- a/mindspore/core/ops/transpose.h +++ b/mindspore/core/ops/transpose.h @@ -18,27 +18,25 @@ #define MINDSPORE_CORE_OPS_TRANSPOSE_H_ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { -constexpr auto kNameTranspose = prim::kTranspose; +constexpr auto kNameTranspose = "Transpose"; /// \brief Permutes the dimensions of the input tensor according to input permutation. /// Refer to Python API @ref mindspore.ops.Transpose for more details. -class MS_CORE_API Transpose : public PrimitiveC { +class MIND_API Transpose : public BaseOperator { public: + MIND_API_BASE_MEMBER(Transpose); /// \brief Constructor. - Transpose() : PrimitiveC(prim::kTranspose) { InitIOName({"x", "perm"}, {"output"}); } - /// \brief Destructor. - ~Transpose() = default; - MS_DECLARE_PARENT(Transpose, PrimitiveC); + Transpose() : BaseOperator(kNameTranspose) { InitIOName({"x", "perm"}, {"output"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr TransposeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TransposeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/trunc.cc b/mindspore/core/ops/trunc.cc index 66187c3367..c5f4179b85 100644 --- a/mindspore/core/ops/trunc.cc +++ b/mindspore/core/ops/trunc.cc @@ -23,6 +23,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -44,6 +45,7 @@ TypePtr TruncInferType(const PrimitivePtr &prim, const std::vector &input_args) { return abstract::MakeAbstract(TruncInferShape(primitive, input_args), TruncInferType(primitive, input_args)); diff --git a/mindspore/core/ops/trunc.h b/mindspore/core/ops/trunc.h index caee3f8337..272cee6f0c 100644 --- a/mindspore/core/ops/trunc.h +++ b/mindspore/core/ops/trunc.h @@ -20,25 +20,23 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTrunc = "Trunc"; /// \brief Returns a new tensor with the truncated integer values of the elements of input. -class Trunc : public PrimitiveC { +class Trunc : public BaseOperator { public: + MIND_API_BASE_MEMBER(Trunc); /// \brief Constructor. - Trunc() : PrimitiveC(kNameTrunc) { InitIOName({"input_x"}, {"output_y"}); } - /// \brief Destructor. - ~Trunc() = default; - MS_DECLARE_PARENT(Trunc, PrimitiveC); + Trunc() : BaseOperator(kNameTrunc) { InitIOName({"input_x"}, {"output_y"}); } }; -AbstractBasePtr TruncInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TruncInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimTruncPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/truncate_div.cc b/mindspore/core/ops/truncate_div.cc index c46bcad4a5..60909944d4 100644 --- a/mindspore/core/ops/truncate_div.cc +++ b/mindspore/core/ops/truncate_div.cc @@ -24,7 +24,7 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" -#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -83,6 +83,7 @@ TypePtr TruncateDivInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/truncate_div.h b/mindspore/core/ops/truncate_div.h index 3742dc4551..d5613455a1 100644 --- a/mindspore/core/ops/truncate_div.h +++ b/mindspore/core/ops/truncate_div.h @@ -17,23 +17,21 @@ #ifndef MINDSPORE_CORE_OPS_TRUNCATE_DIV_H_ #define MINDSPORE_CORE_OPS_TRUNCATE_DIV_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTruncateDiv = "TruncateDiv"; -class TruncateDiv : public PrimitiveC { +class MIND_API TruncateDiv : public BaseOperator { public: - TruncateDiv() : PrimitiveC(kNameTruncateDiv) { InitIOName({"x", "y"}, {"output"}); } - ~TruncateDiv() = default; - MS_DECLARE_PARENT(TruncateDiv, PrimitiveC); + MIND_API_BASE_MEMBER(TruncateDiv); + TruncateDiv() : BaseOperator(kNameTruncateDiv) { InitIOName({"x", "y"}, {"output"}); } void Init() {} }; -AbstractBasePtr TruncateDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TruncateDivInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/truncate_mod.cc b/mindspore/core/ops/truncate_mod.cc index faac5fa855..bd7d87deed 100644 --- a/mindspore/core/ops/truncate_mod.cc +++ b/mindspore/core/ops/truncate_mod.cc @@ -25,6 +25,7 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -62,6 +63,7 @@ TypePtr TruncateModInferType(const PrimitivePtr &prim, const std::vector &input_args) { auto shape = TruncateModInferShape(primitive, input_args); diff --git a/mindspore/core/ops/truncate_mod.h b/mindspore/core/ops/truncate_mod.h index 44d81213c4..4d837ee3c6 100644 --- a/mindspore/core/ops/truncate_mod.h +++ b/mindspore/core/ops/truncate_mod.h @@ -17,23 +17,21 @@ #ifndef MINDSPORE_CORE_OPS_TRUNCATE_MOD_H_ #define MINDSPORE_CORE_OPS_TRUNCATE_MOD_H_ #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameTruncateMod = "TruncateMod"; -class TruncateMod : public PrimitiveC { +class TruncateMod : public BaseOperator { public: - TruncateMod() : PrimitiveC(kNameTruncateMod) { InitIOName({"x", "y"}, {"output"}); } - ~TruncateMod() = default; - MS_DECLARE_PARENT(TruncateMod, PrimitiveC); + MIND_API_BASE_MEMBER(TruncateMod); + TruncateMod() : BaseOperator(kNameTruncateMod) { InitIOName({"x", "y"}, {"output"}); } void Init() {} }; -AbstractBasePtr TruncateModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr TruncateModInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tuple_get_item.cc b/mindspore/core/ops/tuple_get_item.cc index 6eda4f4c3a..a532eb5391 100644 --- a/mindspore/core/ops/tuple_get_item.cc +++ b/mindspore/core/ops/tuple_get_item.cc @@ -15,9 +15,12 @@ */ #include "ops/tuple_get_item.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(TupleGetItem, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameTupleGetItem, TupleGetItem); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/tuple_get_item.h b/mindspore/core/ops/tuple_get_item.h index 9c947827eb..07bcd0b11f 100644 --- a/mindspore/core/ops/tuple_get_item.h +++ b/mindspore/core/ops/tuple_get_item.h @@ -16,22 +16,18 @@ #ifndef MINDSPORE_CORE_OPS_TUPLE_GET_ITEM_H_ #define MINDSPORE_CORE_OPS_TUPLE_GET_ITEM_H_ -#include "ops/primitive_c.h" +#include "ops/base_operator.h" namespace mindspore { namespace ops { constexpr auto kNameTupleGetItem = "TupleGetItem"; /// \brief TupleGetItem op is added to the multi-output node to describe which output of the node, which is only used /// in FuncGraph. -class MS_CORE_API TupleGetItem : public PrimitiveC { +class MIND_API TupleGetItem : public BaseOperator { public: + MIND_API_BASE_MEMBER(TupleGetItem); /// \brief Constructor. - TupleGetItem() : PrimitiveC(kNameTupleGetItem) {} - - /// \brief Destructor. - ~TupleGetItem() = default; - - MS_DECLARE_PARENT(TupleGetItem, PrimitiveC); + TupleGetItem() : BaseOperator(kNameTupleGetItem) {} }; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/uniform_real.cc b/mindspore/core/ops/uniform_real.cc index 00946eb26d..d4cfda9dcf 100644 --- a/mindspore/core/ops/uniform_real.cc +++ b/mindspore/core/ops/uniform_real.cc @@ -19,17 +19,19 @@ #include #include "ops/op_utils.h" #include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(UniformReal, PrimitiveC, BaseOperator); void UniformReal::Init(int64_t seed, int64_t seed2) { this->set_seed(seed); this->set_seed2(seed2); } -void UniformReal::set_seed(int64_t seed) { (void)this->AddAttr(kSeed, MakeValue(seed)); } +void UniformReal::set_seed(int64_t seed) { (void)this->AddAttr(kSeed, api::MakeValue(seed)); } -void UniformReal::set_seed2(int64_t seed2) { (void)this->AddAttr(kSeed2, MakeValue(seed2)); } +void UniformReal::set_seed2(int64_t seed2) { (void)this->AddAttr(kSeed2, api::MakeValue(seed2)); } int64_t UniformReal::get_seed() const { auto value_ptr = GetAttr(kSeed); diff --git a/mindspore/core/ops/uniform_real.h b/mindspore/core/ops/uniform_real.h index 7f5c45853f..379687a7ff 100644 --- a/mindspore/core/ops/uniform_real.h +++ b/mindspore/core/ops/uniform_real.h @@ -20,9 +20,9 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { @@ -30,13 +30,11 @@ constexpr auto kNameUniformReal = "UniformReal"; /// \brief Produces random floating-point values i, uniformly distributed to the interval [0, 1). /// Refer to Python API @ref mindspore.ops.UniformReal for more details. -class MS_CORE_API UniformReal : public PrimitiveC { +class MIND_API UniformReal : public BaseOperator { public: + MIND_API_BASE_MEMBER(UniformReal); /// \brief Constructor. - UniformReal() : PrimitiveC(kNameUniformReal) {} - /// \brief Destructor. - ~UniformReal() = default; - MS_DECLARE_PARENT(UniformReal, PrimitiveC); + UniformReal() : BaseOperator(kNameUniformReal) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.UniformReal for the inputs. void Init(int64_t seed, int64_t seed2); /// \brief Set seed. diff --git a/mindspore/core/ops/unique.cc b/mindspore/core/ops/unique.cc index d51a1cdeba..c22d625118 100644 --- a/mindspore/core/ops/unique.cc +++ b/mindspore/core/ops/unique.cc @@ -15,9 +15,12 @@ */ #include "ops/unique.h" +#include "ops/primitive_c.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Unique, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameUnique, Unique); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/unique.h b/mindspore/core/ops/unique.h index 88753ab963..239f3e8826 100644 --- a/mindspore/core/ops/unique.h +++ b/mindspore/core/ops/unique.h @@ -16,9 +16,8 @@ #ifndef MINDSPORE_CORE_OPS_UNIQUE_H_ #define MINDSPORE_CORE_OPS_UNIQUE_H_ -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { @@ -26,13 +25,11 @@ constexpr auto kNameUnique = "Unique"; /// \brief Returns the unique elements of input tensor and also returns a tensor containing /// the index of each value of input tensor corresponding to the output unique tensor. /// Refer to Python API @ref mindspore.ops.Unique for more details. -class MS_CORE_API Unique : public PrimitiveC { +class MIND_API Unique : public BaseOperator { public: + MIND_API_BASE_MEMBER(Unique); /// \brief Constructor. - Unique() : PrimitiveC(kNameUnique) { InitIOName({"x", "y"}, {"output"}); } - /// \brief Destructor. - ~Unique() = default; - MS_DECLARE_PARENT(Unique, PrimitiveC); + Unique() : BaseOperator(kNameUnique) { InitIOName({"x", "y"}, {"output"}); } /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/unpack.cc b/mindspore/core/ops/unpack.cc index d58fb6d25d..2c7c10cbe5 100644 --- a/mindspore/core/ops/unpack.cc +++ b/mindspore/core/ops/unpack.cc @@ -15,11 +15,16 @@ */ #include "ops/unpack.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Unpack, PrimitiveC, BaseOperator); void Unpack::Init(const int64_t axis) { this->set_axis(axis); } -void Unpack::set_axis(const int64_t axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void Unpack::set_axis(const int64_t axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } int64_t Unpack::get_axis() const { return GetValue(GetAttr(kAxis)); } REGISTER_PRIMITIVE_C(kNameUnpack, Unpack); diff --git a/mindspore/core/ops/unpack.h b/mindspore/core/ops/unpack.h index 9db5d0b207..e6678ad094 100644 --- a/mindspore/core/ops/unpack.h +++ b/mindspore/core/ops/unpack.h @@ -22,23 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameUnpack = "Unpack"; /// \brief Unstacks tensor in specified axis. Refer to Python API @ref mindspore.ops.Unstack for more details. -class MS_CORE_API Unpack : public PrimitiveC { +class MIND_API Unpack : public BaseOperator { public: + MIND_API_BASE_MEMBER(Unpack); /// \brief Constructor. - Unpack() : PrimitiveC(kNameUnpack) {} - /// \brief Destructor. - ~Unpack() = default; - MS_DECLARE_PARENT(Unpack, PrimitiveC); + Unpack() : BaseOperator(kNameUnpack) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Unstack for the inputs. void Init(const int64_t axis = 0); /// \brief Set axis. @@ -48,8 +44,8 @@ class MS_CORE_API Unpack : public PrimitiveC { /// \return axis. int64_t get_axis() const; }; -AbstractBasePtr UnpackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr UnpackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/unsorted_segment_sum.cc b/mindspore/core/ops/unsorted_segment_sum.cc index 53723e8884..7a0a2e08ac 100644 --- a/mindspore/core/ops/unsorted_segment_sum.cc +++ b/mindspore/core/ops/unsorted_segment_sum.cc @@ -22,9 +22,11 @@ #include "ops/op_utils.h" #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(UnsortedSegmentSum, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameUnsortedSegmentSum, UnsortedSegmentSum); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/unsorted_segment_sum.h b/mindspore/core/ops/unsorted_segment_sum.h index cffe15c5a5..61f7d7c4e8 100644 --- a/mindspore/core/ops/unsorted_segment_sum.h +++ b/mindspore/core/ops/unsorted_segment_sum.h @@ -21,30 +21,28 @@ #include #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameUnsortedSegmentSum = "UnsortedSegmentSum"; /// \brief Computes the sum of a tensor along segments. /// Refer to Python API @ref mindspore.ops.UnsortedSegmentSum for more details. -class MS_CORE_API UnsortedSegmentSum : public PrimitiveC { +class MIND_API UnsortedSegmentSum : public BaseOperator { public: + MIND_API_BASE_MEMBER(UnsortedSegmentSum); /// \brief Constructor. - UnsortedSegmentSum() : PrimitiveC(kNameUnsortedSegmentSum) { + UnsortedSegmentSum() : BaseOperator(kNameUnsortedSegmentSum) { InitIOName({"x", "segment_ids", "num_segments"}, {"y"}); } - /// \brief Destructor. - ~UnsortedSegmentSum() = default; - MS_DECLARE_PARENT(UnsortedSegmentSum, PrimitiveC); /// \brief Init. void Init() const {} }; -AbstractBasePtr UnsortedSegmentSumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr UnsortedSegmentSumInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/unsqueeze.cc b/mindspore/core/ops/unsqueeze.cc index 92817d4831..d0d93af7d7 100644 --- a/mindspore/core/ops/unsqueeze.cc +++ b/mindspore/core/ops/unsqueeze.cc @@ -18,12 +18,14 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Unsqueeze, PrimitiveC, BaseOperator); void Unsqueeze::Init(const std::vector axis) { this->set_axis(axis); } -void Unsqueeze::set_axis(const std::vector axis) { (void)this->AddAttr(kAxis, MakeValue(axis)); } +void Unsqueeze::set_axis(const std::vector axis) { (void)this->AddAttr(kAxis, api::MakeValue(axis)); } std::vector Unsqueeze::get_axis() const { return GetValue>(GetAttr(kAxis)); } diff --git a/mindspore/core/ops/unsqueeze.h b/mindspore/core/ops/unsqueeze.h index b952cdd313..98ee54f09d 100644 --- a/mindspore/core/ops/unsqueeze.h +++ b/mindspore/core/ops/unsqueeze.h @@ -20,23 +20,18 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameUnsqueeze = "Unsqueeze"; /// \brief Unsqueeze defined the Unsqueeze operator prototype of lite. -class MS_CORE_API Unsqueeze : public PrimitiveC { +class MIND_API Unsqueeze : public BaseOperator { public: + MIND_API_BASE_MEMBER(Unsqueeze); /// \brief Constructor. - Unsqueeze() : PrimitiveC(kNameUnsqueeze) {} - - /// \brief Destructor. - ~Unsqueeze() = default; - - MS_DECLARE_PARENT(Unsqueeze, PrimitiveC); + Unsqueeze() : BaseOperator(kNameUnsqueeze) {} /// \brief Method to init the op's attributes /// @@ -53,8 +48,8 @@ class MS_CORE_API Unsqueeze : public PrimitiveC { /// \return dimensions info of expanding. std::vector get_axis() const; }; -AbstractBasePtr UnsqueezeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr UnsqueezeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/unstack.cc b/mindspore/core/ops/unstack.cc index 081cc6295a..6574c703fd 100644 --- a/mindspore/core/ops/unstack.cc +++ b/mindspore/core/ops/unstack.cc @@ -15,11 +15,16 @@ */ #include "ops/unstack.h" +#include "utils/check_convert_utils.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Unstack, PrimitiveC, BaseOperator); void Unstack::Init(const int64_t axis) { this->set_axis(axis); } -void Unstack::set_axis(const int64_t axis) { (void)AddAttr(kAxis, MakeValue(axis)); } +void Unstack::set_axis(const int64_t axis) { (void)AddAttr(kAxis, api::MakeValue(axis)); } int64_t Unstack::get_axis() const { return GetValue(GetAttr(kAxis)); } REGISTER_PRIMITIVE_C(kNameUnstack, Unstack); diff --git a/mindspore/core/ops/unstack.h b/mindspore/core/ops/unstack.h index 28846eca87..bd64f30ed6 100644 --- a/mindspore/core/ops/unstack.h +++ b/mindspore/core/ops/unstack.h @@ -22,23 +22,19 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/primitive_infer_map.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameUnstack = "Unstack"; /// \brief Unstacks tensor in specified axis. Refer to Python API @ref mindspore.ops.Unstack for more details. -class MS_CORE_API Unstack : public PrimitiveC { +class MIND_API Unstack : public BaseOperator { public: + MIND_API_BASE_MEMBER(Unstack); /// \brief Constructor. - Unstack() : PrimitiveC(kNameUnstack) {} - /// \brief Destructor. - ~Unstack() = default; - MS_DECLARE_PARENT(Unstack, PrimitiveC); + Unstack() : BaseOperator(kNameUnstack) {} /// \brief Init. Refer to the parameters of Python API @ref mindspore.ops.Unstack for the inputs. void Init(const int64_t axis = 0); /// \brief Set axis. @@ -49,8 +45,8 @@ class MS_CORE_API Unstack : public PrimitiveC { int64_t get_axis() const; }; -AbstractBasePtr UnstackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr UnstackInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/upper_bound.cc b/mindspore/core/ops/upper_bound.cc index 924381f157..90e1f3cf9c 100644 --- a/mindspore/core/ops/upper_bound.cc +++ b/mindspore/core/ops/upper_bound.cc @@ -15,6 +15,10 @@ */ #include "ops/upper_bound.h" +#include "ops/op_utils.h" +#include "abstract/primitive_infer_map.h" +#include "utils/check_convert_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -59,6 +63,7 @@ TypePtr UpperBoundInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/upper_bound.h b/mindspore/core/ops/upper_bound.h index 9c9dae7ef7..78322a445e 100644 --- a/mindspore/core/ops/upper_bound.h +++ b/mindspore/core/ops/upper_bound.h @@ -22,23 +22,19 @@ #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "abstract/primitive_infer_map.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameUpperBound = "UpperBound"; -class UpperBound : public PrimitiveC { +class MIND_API UpperBound : public BaseOperator { public: - UpperBound() : PrimitiveC(kNameUpperBound) { InitIOName({"sorted_x", "values"}, {"y"}); } - ~UpperBound() = default; - MS_DECLARE_PARENT(UpperBound, PrimitiveC); + MIND_API_BASE_MEMBER(UpperBound); + UpperBound() : BaseOperator(kNameUpperBound) { InitIOName({"sorted_x", "values"}, {"y"}); } }; -AbstractBasePtr UpperBoundInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr UpperBoundInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimUpperBound = std::shared_ptr; } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/where.cc b/mindspore/core/ops/where.cc index d01fb7b6f3..ca38fe3af5 100644 --- a/mindspore/core/ops/where.cc +++ b/mindspore/core/ops/where.cc @@ -19,9 +19,11 @@ #include "utils/check_convert_utils.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { +MIND_API_BASE_IMPL(Where, PrimitiveC, BaseOperator); REGISTER_PRIMITIVE_C(kNameWhere, Where); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/where.h b/mindspore/core/ops/where.h index e532e61b76..10808e78ec 100644 --- a/mindspore/core/ops/where.h +++ b/mindspore/core/ops/where.h @@ -19,30 +19,25 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameWhere = "Where"; /// \brief Where defined the operator prototype of selecting values which meet condition. -class MS_CORE_API Where : public PrimitiveC { +class MIND_API Where : public BaseOperator { public: + MIND_API_BASE_MEMBER(Where); /// \brief Constructor. - Where() : PrimitiveC(kNameWhere) { InitIOName({"condition"}, {"output"}); } - - /// \brief Destructor. - ~Where() = default; - - MS_DECLARE_PARENT(Where, PrimitiveC); + Where() : BaseOperator(kNameWhere) { InitIOName({"condition"}, {"output"}); } /// \brief Method to init the op's attributes void Init() const {} }; -AbstractBasePtr WhereInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr WhereInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/xdivy.cc b/mindspore/core/ops/xdivy.cc index 56784fd641..c5c603fbc5 100644 --- a/mindspore/core/ops/xdivy.cc +++ b/mindspore/core/ops/xdivy.cc @@ -26,6 +26,8 @@ #include "abstract/abstract_value.h" #include "abstract/primitive_infer_map.h" #include "ops/op_utils.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -61,6 +63,8 @@ TypePtr XdivyInferType(const PrimitivePtr &primitive, const std::vector &input_args) { auto shape = XdivyInferShape(primitive, input_args); diff --git a/mindspore/core/ops/xdivy.h b/mindspore/core/ops/xdivy.h index a5930cc4f6..bff2b609d3 100644 --- a/mindspore/core/ops/xdivy.h +++ b/mindspore/core/ops/xdivy.h @@ -22,24 +22,22 @@ #include #include #include -#include "ops/primitive_c.h" -#include "ops/op_utils.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameXdivy = "Xdivy"; -class Xdivy : public PrimitiveC { +class Xdivy : public BaseOperator { public: - Xdivy() : PrimitiveC(prim::kPrimXdivy->name()) { InitIOName({"x", "y"}, {"output"}); } - ~Xdivy() = default; - MS_DECLARE_PARENT(Xdivy, PrimitiveC); + MIND_API_BASE_MEMBER(Xdivy); + Xdivy() : BaseOperator("Xdivy") { InitIOName({"x", "y"}, {"output"}); } }; using PrimXdivyPtr = std::shared_ptr; -AbstractBasePtr XdivyInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr XdivyInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/ops/xlogy.cc b/mindspore/core/ops/xlogy.cc index b99f329251..c70e5aa0ed 100644 --- a/mindspore/core/ops/xlogy.cc +++ b/mindspore/core/ops/xlogy.cc @@ -25,6 +25,8 @@ #include "ops/op_utils.h" #include "abstract/abstract_value.h" #include "ops/primitive_c.h" +#include "mindapi/src/helper.h" + namespace mindspore { namespace ops { namespace { @@ -45,6 +47,8 @@ TypePtr XlogyInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/xlogy.h b/mindspore/core/ops/xlogy.h index 4004cc157a..c86d32a682 100644 --- a/mindspore/core/ops/xlogy.h +++ b/mindspore/core/ops/xlogy.h @@ -21,22 +21,20 @@ #include #include #include -#include "ops/op_utils.h" -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { constexpr auto kNameXlogy = "Xlogy"; -class Xlogy : public PrimitiveC { +class MIND_API Xlogy : public BaseOperator { public: - Xlogy() : PrimitiveC(kNameXlogy) { InitIOName({"x", "y"}, {"output"}); } - ~Xlogy() = default; - MS_DECLARE_PARENT(Xlogy, PrimitiveC); + MIND_API_BASE_MEMBER(Xlogy); + Xlogy() : BaseOperator(kNameXlogy) { InitIOName({"x", "y"}, {"output"}); } }; -AbstractBasePtr XlogyInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr XlogyInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); using PrimXlogyPtr = std::shared_ptr; } // namespace ops diff --git a/mindspore/core/ops/zeros.cc b/mindspore/core/ops/zeros.cc index 651524b3bf..4ac198919b 100644 --- a/mindspore/core/ops/zeros.cc +++ b/mindspore/core/ops/zeros.cc @@ -21,6 +21,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -77,6 +78,8 @@ ValuePtr ZerosInferValue(const PrimitivePtr &prim, const std::vector #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" -#include "ops/op_utils.h" + +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Creates a tensor filled with value zeros. Refer to Python API @ref mindspore.ops.Zeros for more details. -class MS_CORE_API Zeros : public PrimitiveC { +class MIND_API Zeros : public BaseOperator { public: + MIND_API_BASE_MEMBER(Zeros); /// \brief Constructor. - Zeros() : PrimitiveC(prim::kPrimZeros->name()) {} - /// \brief Destructor. - ~Zeros() = default; - MS_DECLARE_PARENT(Zeros, PrimitiveC); + Zeros() : BaseOperator("Zeros") {} /// \brief Init. void Init() const {} }; diff --git a/mindspore/core/ops/zeros_like.cc b/mindspore/core/ops/zeros_like.cc index fc5455cf61..648c697c78 100644 --- a/mindspore/core/ops/zeros_like.cc +++ b/mindspore/core/ops/zeros_like.cc @@ -23,6 +23,7 @@ #include "utils/check_convert_utils.h" #include "utils/tensor_construct_utils.h" #include "abstract/primitive_infer_map.h" +#include "mindapi/src/helper.h" namespace mindspore { namespace ops { @@ -41,6 +42,8 @@ TypePtr ZerosLikeInferType(const PrimitivePtr &primitive, const std::vector &input_args) { MS_EXCEPTION_IF_NULL(primitive); diff --git a/mindspore/core/ops/zeros_like.h b/mindspore/core/ops/zeros_like.h index bfc490f23a..10eecb083f 100644 --- a/mindspore/core/ops/zeros_like.h +++ b/mindspore/core/ops/zeros_like.h @@ -19,25 +19,22 @@ #include #include -#include "ops/primitive_c.h" -#include "abstract/abstract_value.h" -#include "utils/check_convert_utils.h" +#include "ops/base_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace ops { /// \brief Creates a new tensor. Refer to Python API @ref mindspore.ops.ZerosLike for more details. -class MS_CORE_API ZerosLike : public PrimitiveC { +class MIND_API ZerosLike : public BaseOperator { public: + MIND_API_BASE_MEMBER(ZerosLike); /// \brief Constructor. - ZerosLike() : PrimitiveC(prim::kPrimZerosLike->name()) { InitIOName({"x"}, {"y"}); } - /// \brief Destructor. - ~ZerosLike() = default; - MS_DECLARE_PARENT(ZerosLike, PrimitiveC); + ZerosLike() : BaseOperator("ZerosLike") { InitIOName({"x"}, {"y"}); } /// \brief Init. void Init() const {} }; -AbstractBasePtr ZerosLikeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, - const std::vector &input_args); +abstract::AbstractBasePtr ZerosLikeInfer(const abstract::AnalysisEnginePtr &, const PrimitivePtr &primitive, + const std::vector &input_args); } // namespace ops } // namespace mindspore diff --git a/mindspore/core/utils/check_convert_utils.h b/mindspore/core/utils/check_convert_utils.h index 817047c775..8ef2d6e56c 100644 --- a/mindspore/core/utils/check_convert_utils.h +++ b/mindspore/core/utils/check_convert_utils.h @@ -69,11 +69,6 @@ enum ReduceType : int64_t { REDUCE_UNKNOW = 7, }; -enum PoolMode : int64_t { - MAX_POOLING = 0, - MEAN_POOLING = 1, -}; - enum GateOrderMode : int64_t { RZH = 0, ZRH = 1 }; template diff --git a/mindspore/lite/build_lite.sh b/mindspore/lite/build_lite.sh index 0739691c19..a393929d33 100755 --- a/mindspore/lite/build_lite.sh +++ b/mindspore/lite/build_lite.sh @@ -263,7 +263,7 @@ build_lite() { CMAKE_TOOLCHAIN_FILE=${BASEPATH}/cmake/lite_ios.cmake fi - BRANCH_NAME=nnie_3516_master_dev + BRANCH_NAME=nnie_3516_master if [[ ("${MSLITE_REGISTRY_DEVICE}" == "Hi3516D" || "${TOOLCHAIN_NAME}" == "himix200") && "${local_lite_platform}" == "arm32" ]]; then TOOLCHAIN_NAME="himix200" MSLITE_REGISTRY_DEVICE=Hi3516D diff --git a/mindspore/lite/include/registry/model_parser.h b/mindspore/lite/include/registry/model_parser.h index e7431b5190..df56ecfe34 100644 --- a/mindspore/lite/include/registry/model_parser.h +++ b/mindspore/lite/include/registry/model_parser.h @@ -17,7 +17,7 @@ #ifndef MINDSPORE_LITE_INCLUDE_REGISTRY_MODEL_PARSER_H_ #define MINDSPORE_LITE_INCLUDE_REGISTRY_MODEL_PARSER_H_ -#include "api/ir/func_graph.h" +#include "mindapi/ir/func_graph.h" #include "include/registry/converter_context.h" namespace mindspore { diff --git a/mindspore/lite/include/registry/node_parser.h b/mindspore/lite/include/registry/node_parser.h index 80439960b0..99f942b830 100644 --- a/mindspore/lite/include/registry/node_parser.h +++ b/mindspore/lite/include/registry/node_parser.h @@ -22,6 +22,7 @@ #include #include #include "include/registry/converter_context.h" +#include "ops/base_operator.h" namespace onnx { class GraphProto; @@ -45,7 +46,7 @@ struct ModelT; namespace mindspore { namespace ops { /// \brief PrimitiveC defined a base class for storing properties -class PrimitiveC; +using BaseOperatorPtr = api::SharedPtr; } // namespace ops namespace converter { /// \brief NodeParser defined a base class for parsing node's attributes. @@ -63,7 +64,7 @@ class MS_API NodeParser { /// \param[in] onnx_node Define the node to be resolved. /// /// \return PrimitiveC Attribute storage. - virtual ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { + virtual ops::BaseOperatorPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { return nullptr; } @@ -73,7 +74,7 @@ class MS_API NodeParser { /// \param[in] weight Define the node which contains weight information. /// /// \return PrimitiveC Attribute storage. - virtual ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + virtual ops::BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { return nullptr; } @@ -86,9 +87,9 @@ class MS_API NodeParser { /// \param[in] output_size Define the output num of current node, which need to be determined by user. /// /// \return PrimitiveC Attribute storage. - virtual ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { + virtual ops::BaseOperatorPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { return nullptr; } @@ -98,9 +99,9 @@ class MS_API NodeParser { /// \param[in] tflite_model Define the model, which contains all information abort the graph. /// /// \return PrimitiveC Attribute storage. - virtual ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { + virtual ops::BaseOperatorPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { return nullptr; } }; diff --git a/mindspore/lite/include/registry/pass_base.h b/mindspore/lite/include/registry/pass_base.h index 46870ce076..85ee4f310e 100644 --- a/mindspore/lite/include/registry/pass_base.h +++ b/mindspore/lite/include/registry/pass_base.h @@ -20,7 +20,7 @@ #include #include #include "include/lite_utils.h" -#include "api/ir/func_graph.h" +#include "mindapi/ir/func_graph.h" namespace mindspore { namespace registry { diff --git a/mindspore/lite/src/ops/ops_utils.cc b/mindspore/lite/src/ops/ops_utils.cc index aac85b0b0d..0adbe0ea46 100644 --- a/mindspore/lite/src/ops/ops_utils.cc +++ b/mindspore/lite/src/ops/ops_utils.cc @@ -17,6 +17,7 @@ #include #include #include "src/ops/ops_utils.h" +#include "mindapi/base/shared_ptr.h" #ifdef PRIMITIVE_WRITEABLE #include "mindspore/core/ir/anf.h" @@ -45,811 +46,820 @@ std::unique_ptr GetPrimitiveT(const AnfNodePtr &node) { } } +template +api::SharedPtr GetOperator(const AnfNodePtr &node) { + auto prim = GetValueNode(node); + if (prim == nullptr) { + return nullptr; + } + return api::MakeShared(prim); +} + std::unique_ptr AbsPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AbsGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ActivationPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ActivationGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AdamPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AdderFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AddFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AddGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AddNPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AllPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ApplyMomentumPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ArgMaxFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ArgMinFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AssertPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AssignPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AssignAddPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AudioSpectrogramPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AvgPoolFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AvgPoolGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BatchNormPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BatchToSpacePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BatchToSpaceNDPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BiasAddPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BiasAddGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BNGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr BroadcastToPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CastPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CeilPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ClipPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ConcatPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ConstantOfShapePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr Conv2DBackpropFilterFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr Conv2DBackpropInputFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr Conv2DFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr Conv2dTransposeFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CosPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CropPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CropAndResizePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CustomExtractFeaturesPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CustomNormalizePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CustomPredictPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DependPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DepthToSpacePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DetectionPostProcessPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DivFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DivGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DropoutPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DropoutGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr GRUPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr EltwisePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr EluPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr EmbeddingLookupFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr EqualPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ExpandDimsPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ExpFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FftImagPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FftRealPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FillPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FlattenPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FlattenGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FloorPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FloorDivPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FloorModPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FullConnectionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr FusedBatchNormPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr GatherPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr GatherNdPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr GreaterPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr GreaterEqualPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr HashtableLookupPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr InstanceNormPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr InvertPermutationPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LayerNormFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LayerNormGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LeakyReluPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LessPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LessEqualPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LogPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LogGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LogicalAndPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LogicalNotPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LogicalOrPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LrnPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LpNormalizationPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LshProjectionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LSTMPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LSTMGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LSTMGradDataPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LSTMGradWeightPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr L2NormalizeFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MatMulFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MaximumPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MaximumGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MaxPoolFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MaxPoolGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SwitchLayerPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MfccPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MinimumPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MinimumGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ModPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MulFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr MulGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr NegPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr NegGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr NotEqualPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr NonMaxSuppressionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr OneHotPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr OnesLikePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr PadFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr PartialFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr PowerGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr PowFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr PReLUFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr QuantDTypeCastPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RaggedRangePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RangePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RandomStandardNormalPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RankPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RealDivPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ReciprocalPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ReduceFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ReshapePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ResizePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ResizeGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ReverseV2PrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ReverseSequencePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RfftPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ROIPoolingPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RoundPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RsqrtPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RsqrtGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ScaleFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ScatterNdPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SelectPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SGDPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ShapePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SigmoidCrossEntropyWithLogitsPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SigmoidCrossEntropyWithLogitsGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SinPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SizePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SkipGramPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SliceFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SmoothL1LossPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SmoothL1LossGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SoftmaxPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SoftmaxCrossEntropyWithLogitsPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SpaceToBatchPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SpaceToBatchNDPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SpaceToDepthPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SparseSoftmaxCrossEntropyWithLogitsPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SparseToDensePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SplitPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SqrtPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SqrtGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SquarePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SquaredDifferencePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SqueezePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr StackPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr StridedSlicePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr StridedSliceGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SubFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SubGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SwitchPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorListFromTensorPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorListGetItemPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorListReservePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorListSetItemPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorListStackPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TileFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TopKFusionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TransposePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr UniquePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr UnstackPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr UnsortedSegmentSumPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr UnsqueezePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr WherePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ZerosLikePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ErfPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SplicePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr LogSoftmaxPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CallPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CumSumPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr SplitWithOverlapPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr GluPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorArrayPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorArrayReadPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr TensorArrayWritePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AffinePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AttentionPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ScatterNdUpdatePrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr AllGatherPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr ReduceScatterPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr DynamicQuantPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr RandomNormalPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr NLLLossPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr NLLLossGradPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); return ms_primc != nullptr ? ops::MSOp2SchemaOp(ms_primc.get()) : nullptr; } std::unique_ptr CustomPrimitiveCreator(const AnfNodePtr &node) { - auto ms_primc = GetValueNode>(node); + auto ms_primc = GetOperator(node); auto schema_op = std::make_unique(); if (schema_op == nullptr) { return nullptr; diff --git a/mindspore/lite/src/ops/ops_utils.h b/mindspore/lite/src/ops/ops_utils.h index c4cf5b8b86..c345a16311 100644 --- a/mindspore/lite/src/ops/ops_utils.h +++ b/mindspore/lite/src/ops/ops_utils.h @@ -23,6 +23,8 @@ #include "src/ops/ops_func_declare.h" #ifdef PRIMITIVE_WRITEABLE +#include "abstract/primitive_infer_map.h" + namespace mindspore { namespace lite { typedef std::unique_ptr (*PrimitiveTCreator)(const AnfNodePtr &node); diff --git a/mindspore/lite/test/st/converter_test.cc b/mindspore/lite/test/st/converter_test.cc index b02729f8cf..fcc98e05be 100644 --- a/mindspore/lite/test/st/converter_test.cc +++ b/mindspore/lite/test/st/converter_test.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include #include #include "tools/converter/converter.h" diff --git a/mindspore/lite/test/st/delegate_test.cc b/mindspore/lite/test/st/delegate_test.cc index 52d3e06295..93097d603e 100644 --- a/mindspore/lite/test/st/delegate_test.cc +++ b/mindspore/lite/test/st/delegate_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "gtest/gtest.h" #include "common/common_test.h" #include "include/errorcode.h" diff --git a/mindspore/lite/test/st/graph_test.cc b/mindspore/lite/test/st/graph_test.cc index 29b84b6062..628f9cb701 100644 --- a/mindspore/lite/test/st/graph_test.cc +++ b/mindspore/lite/test/st/graph_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "gtest/gtest.h" #include "common/common_test.h" #include "include/errorcode.h" diff --git a/mindspore/lite/test/st/mindrt_parallel_test.cc b/mindspore/lite/test/st/mindrt_parallel_test.cc index 4c609691ba..6ab5957f1a 100644 --- a/mindspore/lite/test/st/mindrt_parallel_test.cc +++ b/mindspore/lite/test/st/mindrt_parallel_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "gtest/gtest.h" #include "common/common_test.h" #include "include/errorcode.h" diff --git a/mindspore/lite/test/st/scripts/nnie/run_converter_nnie.sh b/mindspore/lite/test/st/scripts/nnie/run_converter_nnie.sh index 6436bcc2c8..87fb682d4b 100755 --- a/mindspore/lite/test/st/scripts/nnie/run_converter_nnie.sh +++ b/mindspore/lite/test/st/scripts/nnie/run_converter_nnie.sh @@ -7,21 +7,13 @@ function Run_Converter() { cd ${x86_path} || exit 1 tar -zxf mindspore-lite-${version}-linux-x64.tar.gz || exit 1 cd ${x86_path}/mindspore-lite-${version}-linux-x64/ || exit 1 + # generate converter_lite config file + ms_config_file=${x86_path}/converter.cfg cp tools/converter/converter/converter_lite ./ || exit 1 export LD_LIBRARY_PATH=${LD_LIBRARY_PATH}:./tools/converter/lib/:./tools/converter/third_party/glog/lib:./tools/converter/providers/Hi3516D/third_party/opencv-4.2.0:./tools/converter/providers/Hi3516D/third_party/protobuf-3.9.0 - export NNIE_MAPPER_PATH=./tools/converter/providers/Hi3516D/libnnie_mapper.so - export NNIE_DATA_PROCESS_PATH=./tools/converter/providers/Hi3516D/libmslite_nnie_data_process.so export LD_LIBRARY_PATH=${LD_LIBRARY_PATH}:./runtime/lib/ - export BENCHMARK_PATH=${x86_path}/mindspore-lite-${version}-linux-x64/tools/benchmark/benchmark - export NNIE_MODEL_NAME=model.ms - - # generate converter_lite config file - ms_config_file=${x86_path}/converter_for_nnie.cfg - echo "[registry]" > ${ms_config_file} - echo 'plugin_path='${x86_path}'/mindspore-lite-'${version}'-linux-x64/tools/converter/providers/Hi3516D/libmslite_nnie_converter.so' >> ${ms_config_file} - echo -e 'disable_fusion=off\n' >> ${ms_config_file} echo ' ' > ${run_converter_log_file} rm -rf ${ms_models_path} @@ -37,8 +29,15 @@ function Run_Converter() { model_info=`echo ${nnie_line_info}|awk -F ' ' '{print $2}'` model_name=${model_info%%;*} cp ${models_path}/${model_location}/${model_name}.cfg ./ || exit 1 - echo 'export NNIE_CONFIG_PATH=./'${model_name}'.cfg' >> "${run_converter_log_file}" - export NNIE_CONFIG_PATH=./${model_name}.cfg + echo "[registry]" > ${ms_config_file} + echo 'plugin_path='${x86_path}'/mindspore-lite-'${version}'-linux-x64/tools/converter/providers/Hi3516D/libmslite_nnie_converter.so' >> ${ms_config_file} + echo '[nnie]' >> ${ms_config_file} + echo 'nnie_mapper_path=./tools/converter/providers/Hi3516D/libnnie_mapper.so' >> ${ms_config_file} + echo 'nnie_data_process_path=./tools/converter/providers/Hi3516D/libmslite_nnie_data_process.so' >> ${ms_config_file} + echo 'benchmark_path='${x86_path}'/mindspore-lite-'${version}'-linux-x64/tools/benchmark/benchmark' >> ${ms_config_file} + echo 'nnie_config_path='./${model_name}.cfg >> ${ms_config_file} + echo -e 'nnie_disable_inplace_fusion=off\n' >> ${ms_config_file} + echo ${model_name} >> "${run_converter_log_file}" echo './converter_lite --fmk=CAFFE --modelFile='${models_path}'/'${model_location}'/model/'${model_name}'.prototxt --weightFile='${models_path}'/'${model_location}'/model/'${model_name}'.caffemodel --configFile='${ms_config_file}' --outputFile='${ms_models_path}'/'${model_name}'' >> "${run_converter_log_file}" ./converter_lite --fmk=CAFFE --modelFile=${models_path}/${model_location}/model/${model_name}.prototxt --weightFile=${models_path}/${model_location}/model/${model_name}.caffemodel --configFile=${ms_config_file} --outputFile=${ms_models_path}/${model_name} diff --git a/mindspore/lite/test/ut/tools/converter/registry/model_parser_registry_test.cc b/mindspore/lite/test/ut/tools/converter/registry/model_parser_registry_test.cc index 0ae2ae02e0..c1f74b6139 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/model_parser_registry_test.cc +++ b/mindspore/lite/test/ut/tools/converter/registry/model_parser_registry_test.cc @@ -18,10 +18,19 @@ #include "common/common_test.h" #include "ut/tools/converter/registry/parser/model_parser_test.h" #include "tools/optimizer/common/gllo_utils.h" +#include "mindspore/core/ir/anf.h" +#include "mindapi/ir/func_graph.h" using mindspore::converter::ConverterParameters; using mindspore::converter::kFmkTypeCaffe; namespace mindspore { +namespace { +FuncGraphPtr ConvertGraph(api::FuncGraphPtr func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} +} // namespace + class ModelParserRegistryTest : public mindspore::CommonTest { public: ModelParserRegistryTest() = default; @@ -40,7 +49,8 @@ TEST_F(ModelParserRegistryTest, TestRegistry) { ConverterParameters converter_parameters; auto func_graph = model_parser->Parse(converter_parameters); ASSERT_NE(func_graph, nullptr); - auto node_list = func_graph->TopoSort(func_graph->get_return()); + auto graph = ConvertGraph(func_graph); + auto node_list = graph->TopoSort(graph->get_return()); std::vector cnode_list; for (auto &node : node_list) { if (node->isa()) { diff --git a/mindspore/lite/test/ut/tools/converter/registry/node_parser_registry_test.cc b/mindspore/lite/test/ut/tools/converter/registry/node_parser_registry_test.cc index 41010efcb4..455c278828 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/node_parser_registry_test.cc +++ b/mindspore/lite/test/ut/tools/converter/registry/node_parser_registry_test.cc @@ -14,23 +14,27 @@ * limitations under the License. */ -#include "api/ir/func_graph.h" +#define USE_DEPRECATED_API +#include "include/registry/node_parser_registry.h" #include "common/common_test.h" #include "include/registry/model_parser.h" #include "include/registry/model_parser_registry.h" -#include "include/registry/node_parser_registry.h" +#include "mindapi/ir/func_graph.h" +#include "mindspore/core/ir/anf.h" +#include "mindspore/core/ir/func_graph.h" #include "ops/addn.h" #include "proto/graph.pb.h" using mindspore::converter::kFmkTypeTf; +using PrimitiveCPtr = std::shared_ptr; namespace mindspore { namespace converter { class AddNodeParser : public NodeParser { public: - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override { - auto prim = std::make_unique(); + ops::BaseOperatorPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override { + auto prim = api::MakeShared(); if (prim == nullptr) { MS_LOG(ERROR) << "make a shared_ptr failed."; return nullptr; @@ -39,7 +43,7 @@ class AddNodeParser : public NodeParser { for (int i = 0; i < tf_op.input_size(); ++i) { inputs->push_back(tf_op.input(i)); } - return prim.release(); + return prim; } }; REG_NODE_PARSER(kFmkTypeTf, Add, std::make_shared()); @@ -63,17 +67,17 @@ class NodeParserRegistryTest : public CommonTest { TEST_F(NodeParserRegistryTest, TestRegistry) { ASSERT_NE(func_graph_, nullptr); - auto node_list = api::FuncGraph::TopoSort(func_graph_->get_return()); - std::vector cnodes; + auto node_list = mindspore::api::FuncGraph::TopoSort(func_graph_->get_return()); + std::vector cnodes; for (auto &node : node_list) { - if (node->isa()) { - cnodes.push_back(node->cast()); + if (node->isa()) { + cnodes.push_back(node->cast()); } } ASSERT_EQ(cnodes.size(), 2); auto cnode = cnodes.front(); ASSERT_EQ(cnode->size(), 3); - auto prim = GetValueNode>(cnode->input(0)); + auto prim = api::GetValueNode(cnode->input(0)); ASSERT_NE(prim, nullptr); } } // namespace mindspore diff --git a/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.cc b/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.cc index 8fa8332d0f..7e2245fb9d 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.cc +++ b/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.cc @@ -19,6 +19,11 @@ #include #include "include/errorcode.h" #include "include/registry/model_parser_registry.h" +#include "mindapi/ir/func_graph.h" +#include "mindapi/base/shared_ptr.h" +#include "mindapi/base/type_id.h" +#include "mindapi/ir/tensor.h" +#include "ops/return.h" namespace mindspore { api::FuncGraphPtr ModelParserTest::Parse(const converter::ConverterParameters &flag) { @@ -69,7 +74,7 @@ int ModelParserTest::BuildGraphInputs() { return lite::RET_ERROR; } ShapeVector shape{10, 10}; - auto tensor_info = std::make_shared(TypeId::kNumberTypeFloat32, shape); + auto tensor_info = mindspore::api::MakeShared(kNumberTypeFloat32, shape); if (tensor_info == nullptr) { return lite::RET_ERROR; } @@ -106,7 +111,7 @@ int ModelParserTest::BuildGraphNodes() { MS_LOG(ERROR) << "node parser failed."; return lite::RET_ERROR; } - std::vector anf_inputs; + std::vector anf_inputs; for (auto &input : node_inputs) { if (nodes_.find(input) != nodes_.end()) { anf_inputs.push_back(nodes_[input]); @@ -117,9 +122,9 @@ int ModelParserTest::BuildGraphNodes() { return lite::RET_ERROR; } ShapeVector shape{10, 10}; - auto tensor_info = std::make_shared(TypeId::kNumberTypeFloat32, shape); + auto tensor_info = mindspore::api::MakeShared(kNumberTypeFloat32, shape); auto size = tensor_info->Size(); - memset_s(tensor_info->data_c(), size, 0, size); + memset_s(tensor_info->data(), size, 0, size); parameter->set_abstract(tensor_info->ToAbstract()); parameter->set_default_param(tensor_info); parameter->set_name(input); @@ -127,9 +132,9 @@ int ModelParserTest::BuildGraphNodes() { nodes_.insert(std::make_pair(input, parameter)); } } - auto cnode = res_graph_->NewCNode(std::shared_ptr(primc), anf_inputs); + auto cnode = res_graph_->NewCNode(primc, anf_inputs); cnode->set_fullname_with_scope(node_name); - auto tensor_info = std::make_shared(TypeId::kNumberTypeFloat32, ShapeVector{}); + auto tensor_info = mindspore::api::MakeShared(kNumberTypeFloat32, ShapeVector{}); cnode->set_abstract(tensor_info->ToAbstract()); nodes_.insert(std::make_pair(node_name, cnode)); } @@ -152,7 +157,7 @@ int ModelParserTest::BuildGraphOutputs() { if (nodes_.find(outputs[0]) == nodes_.end()) { return lite::RET_ERROR; } - auto return_prim = std::make_shared("Return"); + auto return_prim = mindspore::api::MakeShared(); auto return_cnode = res_graph_->NewCNode(return_prim, {nodes_[outputs[0]]}); return_cnode->set_fullname_with_scope("Return"); res_graph_->set_return(return_cnode); diff --git a/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.h b/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.h index 3e3ca02343..b62727913a 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.h +++ b/mindspore/lite/test/ut/tools/converter/registry/parser/model_parser_test.h @@ -22,6 +22,7 @@ #include #include "include/registry/model_parser.h" #include "include/registry/model_parser_registry.h" +#include "mindapi/ir/anf.h" #include "ut/tools/converter/registry/parser/node_parser_test.h" namespace mindspore { @@ -35,7 +36,7 @@ class ModelParserTest : public converter::ModelParser { int BuildGraphInputs(); int BuildGraphNodes(); int BuildGraphOutputs(); - std::map nodes_; + std::map nodes_; std::map> model_layers_info_; std::vector model_structure_; }; diff --git a/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.cc b/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.cc index 63499708a2..dd1d61d126 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.cc +++ b/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "ut/tools/converter/registry/parser/node_parser_test.h" #include #include @@ -28,9 +29,9 @@ class AddNodeParserTest : public NodeParserTest { public: AddNodeParserTest() = default; ~AddNodeParserTest() = default; - ops::PrimitiveC *Parse() override { - auto primc = std::make_unique(); - return primc.release(); + BaseOperatorPtr Parse() override { + auto primc = api::MakeShared(); + return primc; } }; @@ -38,11 +39,11 @@ class SplitNodeParserTest : public NodeParserTest { public: SplitNodeParserTest() = default; ~SplitNodeParserTest() = default; - ops::PrimitiveC *Parse() override { - auto primc = std::make_unique(); + BaseOperatorPtr Parse() override { + auto primc = api::MakeShared(); primc->set_axis(0); primc->set_output_num(2); - return primc.release(); + return primc; } }; @@ -50,10 +51,10 @@ class ConcatNodeParserTest : public NodeParserTest { public: ConcatNodeParserTest() = default; ~ConcatNodeParserTest() = default; - ops::PrimitiveC *Parse() override { - auto primc = std::make_unique(); + BaseOperatorPtr Parse() override { + auto primc = api::MakeShared(); primc->set_axis(0); - return primc.release(); + return primc; } }; @@ -62,8 +63,8 @@ class CustomProposalNodeParserTest : public NodeParserTest { public: CustomProposalNodeParserTest() = default; ~CustomProposalNodeParserTest() = default; - ops::PrimitiveC *Parse() override { - auto primc = std::make_unique(); + BaseOperatorPtr Parse() override { + auto primc = api::MakeShared(); primc->set_type("Proposal"); std::map> custom_attrs; std::string height = std::to_string(100); @@ -73,7 +74,7 @@ class CustomProposalNodeParserTest : public NodeParserTest { std::vector width_attr(width.begin(), width.end()); custom_attrs["image_width"] = width_attr; primc->set_attr(custom_attrs); - return primc.release(); + return primc; } }; diff --git a/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.h b/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.h index c435b4d372..7560e852b6 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.h +++ b/mindspore/lite/test/ut/tools/converter/registry/parser/node_parser_test.h @@ -20,17 +20,19 @@ #include #include #include -#include "ops/primitive_c.h" +#include "ops/base_operator.h" +#include "mindapi/base/shared_ptr.h" #include "src/common/log_adapter.h" namespace mindspore { +using BaseOperatorPtr = api::SharedPtr; class NodeParserTest { public: NodeParserTest() = default; virtual ~NodeParserTest() {} - virtual ops::PrimitiveC *Parse() { return nullptr; } + virtual BaseOperatorPtr Parse() { return nullptr; } }; using NodeParserTestPtr = std::shared_ptr; diff --git a/mindspore/lite/test/ut/tools/converter/registry/pass_registry_test.cc b/mindspore/lite/test/ut/tools/converter/registry/pass_registry_test.cc index 44a472a6c8..8d9461c09e 100644 --- a/mindspore/lite/test/ut/tools/converter/registry/pass_registry_test.cc +++ b/mindspore/lite/test/ut/tools/converter/registry/pass_registry_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include #include @@ -27,11 +28,18 @@ #include "tools/converter/optimizer_manager.h" #include "tools/optimizer/common/gllo_utils.h" #include "ut/tools/converter/registry/parser/model_parser_test.h" +#include "ops/op_utils.h" using mindspore::converter::ConverterParameters; using mindspore::converter::kFmkTypeCaffe; using mindspore::registry::POSITION_BEGIN; namespace mindspore { +namespace { +FuncGraphPtr ConvertGraph(api::FuncGraphPtr func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} +} // namespace class PassRegistryTest : public mindspore::CommonTest { public: PassRegistryTest() = default; @@ -59,7 +67,7 @@ class Test1Fusion : public registry::PassBase { if (!opt::CheckPrimitiveType(cnode, prim::kPrimAddFusion)) { return false; } - auto primc = GetValueNode>(cnode->input(0)); + auto primc = mindspore::ops::GetOperator(cnode->input(0)); if (primc == nullptr) { return false; } @@ -76,7 +84,7 @@ class Test1Fusion : public registry::PassBase { return false; } auto input_cnode = input->cast(); - auto add_primc = GetValueNode>(input_cnode->input(0)); + auto add_primc = mindspore::ops::GetOperator(input_cnode->input(0)); if (add_primc == nullptr) { return false; } @@ -94,11 +102,12 @@ class Test1Fusion : public registry::PassBase { if (func_graph == nullptr) { return false; } - auto manager = api::FuncGraphManager::Manage(func_graph); + auto graph = ConvertGraph(func_graph); + auto manager = Manage(graph); if (manager == nullptr) { return false; } - auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); + auto node_list = TopoSort(graph->get_return()); for (auto &node : node_list) { if (!utils::isa(node)) { continue; @@ -120,7 +129,8 @@ class Test1Fusion : public registry::PassBase { } } auto primc = std::make_shared(); - auto new_cnode = func_graph->NewCNode(primc, inputs); + auto primc_c = primc->GetPrim(); + auto new_cnode = graph->NewCNode(primc_c, inputs); new_cnode->set_fullname_with_scope(cnode->fullname_with_scope()); new_cnode->set_abstract(cnode->abstract()->Clone()); manager->Replace(node, new_cnode); @@ -133,7 +143,7 @@ class Test1Fusion : public registry::PassBase { class Test2Fusion : public registry::PassBase { public: Test2Fusion() : PassBase("Test2Fusion") {} - AnfNodePtr CreateCustomOp(const api::FuncGraphPtr func_graph, const CNodePtr &cnode) { + AnfNodePtr CreateCustomOp(const FuncGraphPtr func_graph, const CNodePtr &cnode) { if (func_graph == nullptr || cnode == nullptr) { return nullptr; } @@ -141,6 +151,7 @@ class Test2Fusion : public registry::PassBase { if (primc == nullptr) { return nullptr; } + auto primc_c = primc->GetPrim(); primc->set_type("Custom_AddN"); std::map> custom_attrs; std::string input_num = std::to_string(cnode->size() - 1); @@ -152,7 +163,7 @@ class Test2Fusion : public registry::PassBase { primc->set_attr(custom_attrs); auto inputs = cnode->inputs(); inputs.erase(inputs.begin()); - auto custom_cnode = func_graph->NewCNode(primc, inputs); + auto custom_cnode = func_graph->NewCNode(primc_c, inputs); custom_cnode->set_fullname_with_scope(cnode->fullname_with_scope()); custom_cnode->set_abstract(cnode->abstract()->Clone()); return custom_cnode; @@ -162,11 +173,12 @@ class Test2Fusion : public registry::PassBase { if (func_graph == nullptr) { return false; } - auto manager = api::FuncGraphManager::Manage(func_graph); + auto graph = ConvertGraph(func_graph); + auto manager = Manage(graph); if (manager == nullptr) { return false; } - auto node_list = TopoSort(func_graph->get_return()); + auto node_list = TopoSort(graph->get_return()); for (auto &node : node_list) { if (!utils::isa(node)) { continue; @@ -175,7 +187,7 @@ class Test2Fusion : public registry::PassBase { continue; } auto cnode = node->cast(); - auto custome_cnode = CreateCustomOp(func_graph, cnode); + auto custome_cnode = CreateCustomOp(graph, cnode); if (custome_cnode == nullptr) { return false; } @@ -207,7 +219,8 @@ TEST_F(PassRegistryTest, TestRegistry) { ASSERT_EQ(ret, true); } std::vector cnode_list; - auto node_list = api::FuncGraph::TopoSort(func_graph_->get_return()); + auto graph = ConvertGraph(func_graph_); + auto node_list = TopoSort(graph->get_return()); for (auto &node : node_list) { ASSERT_NE(node, nullptr); if (node->isa()) { @@ -217,7 +230,7 @@ TEST_F(PassRegistryTest, TestRegistry) { ASSERT_EQ(cnode_list.size(), 2); bool is_custom = opt::CheckPrimitiveType(cnode_list.front(), prim::kPrimCustom); ASSERT_EQ(is_custom, true); - auto custome_prim = GetValueNode>(cnode_list.front()->input(0)); + auto custome_prim = mindspore::ops::GetOperator(cnode_list.front()->input(0)); ASSERT_NE(custome_prim, nullptr); auto type = custome_prim->get_type(); ASSERT_EQ(type, std::string("Custom_AddN")); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/activation_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/activation_fusion_test.cc index 0f510fe045..a6cf820ba4 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/activation_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/activation_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/add_concat_act_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/add_concat_act_fusion_test.cc index f24eb8d4ac..5f889a0227 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/add_concat_act_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/add_concat_act_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/constant_folding_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/constant_folding_fusion_test.cc index a4d1f433a2..31684f2bcc 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/constant_folding_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/constant_folding_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_activation_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_activation_fusion_test.cc index de12d84f3a..bcd7bfec63 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_activation_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_activation_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_biasadd_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_biasadd_fusion_test.cc index 5fe1c1a9ed..b7bfcb896e 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_biasadd_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_biasadd_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_bn_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_bn_fusion_test.cc index ee4c8cb8bd..9a0ef2cdf3 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_bn_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_bn_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_scale_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_scale_fusion_test.cc index 0ef7b4382e..b4235adcd4 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/conv_scale_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/conv_scale_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/activation_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/activation_fusion_inout_test.cc index c0cf764d5e..d43194720c 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/activation_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/activation_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "tools/optimizer/fusion/activation_fusion.h" #include "test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.h" @@ -62,13 +63,15 @@ class ActivationFusionInoutTest : public FusionInoutTest { const std::string &name, const float &min_val = FLT_MAX, const float &max_val = -FLT_MAX) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create Act primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(); prim->set_activation_type(type); if (type == ActivationType::HARD_TANH) { prim->set_min_val(min_val); prim->set_max_val(max_val); } - auto act_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto act_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(act_primitive != nullptr, nullptr); auto act = graph->NewCNode({act_primitive, input}); MS_CHECK_TRUE_MSG(act != nullptr, nullptr, "create Act failed"); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/add_concat_act_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/add_concat_act_fusion_inout_test.cc index b0917c7163..2495ffcb85 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/add_concat_act_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/add_concat_act_fusion_inout_test.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include #include #include "test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.h" @@ -29,7 +31,7 @@ namespace mindspore { namespace { constexpr size_t kAddInputTensorWSize = 128; -} +} // namespace class ConcatActFusionInoutTest : public FusionInoutTest { public: ConcatActFusionInoutTest() = default; @@ -75,8 +77,10 @@ class ConcatActFusionInoutTest : public FusionInoutTest { AddParameter(graph_, 0, {add_left_h_, add_left_w_}, kNumberTypeFloat32, "graph_" + name + "_input2"); auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create AddFusion primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(ActivationType::NO_ACTIVATION); - auto add_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto add_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(add_primitive != nullptr, nullptr); auto add_fusion = graph->NewCNode({add_primitive, input1, input2}); MS_CHECK_TRUE_MSG(add_fusion != nullptr, nullptr, "create AddFusion failed"); @@ -88,9 +92,11 @@ class ConcatActFusionInoutTest : public FusionInoutTest { const std::string &name) { auto concat_primitive = std::make_unique(); MS_CHECK_TRUE_MSG(concat_primitive != nullptr, nullptr, "create concat primitivec failed"); + auto prim_c = concat_primitive->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); concat_primitive->Init(); concat_primitive->set_axis(1); - auto concat_primc = NewValueNode(std::shared_ptr(concat_primitive.release())); + auto concat_primc = NewValueNode(prim_c); MS_CHECK_TRUE_RET(concat_primc != nullptr, nullptr); auto concat = graph->NewCNode({concat_primc, input1, input2}); MS_CHECK_TRUE_MSG(concat != nullptr, nullptr, "create Concat failed"); @@ -101,9 +107,11 @@ class ConcatActFusionInoutTest : public FusionInoutTest { CNodePtr AddAct(const FuncGraphPtr &graph, const AnfNodePtr &input, const std::string &name) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create Act primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(); prim->set_activation_type(ActivationType::RELU6); - auto act_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto act_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(act_primitive != nullptr, nullptr); auto act = graph->NewCNode({act_primitive, input}); MS_CHECK_TRUE_MSG(act != nullptr, nullptr, "create Act failed"); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_act_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_act_fusion_inout_test.cc index 82f9e2afe7..34ec02ea21 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_act_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_act_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "tools/optimizer/fusion/conv_activation_fusion.h" #include "test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.h" @@ -57,9 +58,11 @@ class ConvActFusionInoutTest : public ConvFusionInoutTest { static CNodePtr AddAct(const FuncGraphPtr &graph, const AnfNodePtr &input, const std::string &name) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create Act primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(); prim->set_activation_type(ActivationType::RELU); - auto act_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto act_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(act_primitive != nullptr, nullptr); auto act = graph->NewCNode({act_primitive, input}); MS_CHECK_TRUE_MSG(act != nullptr, nullptr, "create Act failed"); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_bias_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_bias_fusion_inout_test.cc index 132f0954b3..9cee238a4e 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_bias_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_bias_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "tools/optimizer/fusion/conv_biasadd_fusion.h" #include "test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.h" @@ -57,8 +58,10 @@ class ConvBiasFusionInoutTest : public ConvFusionInoutTest { static CNodePtr AddBias(const FuncGraphPtr &graph, const AnfNodePtr &input, const std::string &name) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create BiasAdd primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(); - auto bias_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto bias_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(bias_primitive != nullptr, nullptr); auto bias = AddParameter(graph, oc_ * sizeof(float), {oc_}, kNumberTypeFloat32, name + "_bias"); auto bias_add = graph->NewCNode({bias_primitive, input, bias}); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.cc index 8d7f3dd0f3..9e1158cfc3 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.h" #include #include "src/common/log_adapter.h" @@ -25,9 +26,11 @@ namespace mindspore { ValueNodePtr ConvFusionInoutTest::CreateConvPrimitiveValue() { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create Conv2d primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(ic_, oc_, {kh_, kw_}); prim->set_pad_mode(PadMode::SAME); - return NewValueNode(std::shared_ptr(prim.release())); + return NewValueNode(prim_c); } CNodePtr ConvFusionInoutTest::AddConv(const FuncGraphPtr &graph, const AnfNodePtr &input, const std::string &name) { diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.cc index e8c0ff17f5..38e35df611 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.h" #include #include "src/common/log_adapter.h" @@ -101,7 +102,9 @@ CNodePtr FusionInoutTest::AddReturn(const FuncGraphPtr &graph, const std::vector MS_LOG(ERROR) << "new MakeTuple failed"; return nullptr; } - auto return_input_cnode = graph->NewCNode(make_tuple_prim_ptr, return_inputs); + auto prim_c = make_tuple_prim_ptr->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); + auto return_input_cnode = graph->NewCNode(prim_c, return_inputs); if (return_input_cnode == nullptr) { MS_LOG(ERROR) << "new make tuple cnode failed"; return nullptr; @@ -112,7 +115,9 @@ CNodePtr FusionInoutTest::AddReturn(const FuncGraphPtr &graph, const std::vector auto return_prim = std::make_shared(); MS_CHECK_TRUE_MSG(return_prim != nullptr, nullptr, "create return primitivec failed"); - auto return_cnode = graph->NewCNode(return_prim, {return_input}); + auto return_prim_c = return_prim->GetPrim(); + MS_CHECK_TRUE_MSG(return_prim_c != nullptr, nullptr, "prim_c is nullptr"); + auto return_cnode = graph->NewCNode(return_prim_c, {return_input}); MS_CHECK_TRUE_MSG(return_cnode != nullptr, nullptr, "create Return failed"); return_cnode->set_fullname_with_scope("Return"); graph->set_return(return_cnode); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_act_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_act_fusion_inout_test.cc index 4cf4a3a5fc..12e146bd82 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_act_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_act_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "tools/optimizer/fusion/matmul_activation_fusion.h" #include "test/ut/tools/optimizer/fusion/fusion_inout_test/fusion_inout_test.h" @@ -54,9 +55,11 @@ class MatMulActivationFusionInoutTest : public FusionInoutTest { CNodePtr AddAct(const FuncGraphPtr &graph, const AnfNodePtr &input, const std::string &name) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create Act primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(); prim->set_activation_type(ActivationType::RELU); - auto act_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto act_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(act_primitive != nullptr, nullptr); auto act = graph->NewCNode({act_primitive, input}); MS_CHECK_TRUE_MSG(act != nullptr, nullptr, "create Act failed"); @@ -69,8 +72,10 @@ class MatMulActivationFusionInoutTest : public FusionInoutTest { AnfNodePtr bias = AddParameter(graph_, 0, {in_}, kNumberTypeFloat32, "graph_" + name + "_input3"); auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create MatMul primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(false, false, ActivationType::NO_ACTIVATION); - auto matmul_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto matmul_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(matmul_primitive != nullptr, nullptr); auto matmul_fusion = graph->NewCNode({matmul_primitive, input1, input2, bias}); MS_CHECK_TRUE_MSG(matmul_fusion != nullptr, nullptr, "create matmul fusion failed"); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_fusion_inout_test.cc index 2a6e74d0d3..85c69223e8 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_fusion_inout_test.h" #include #include "src/common/log_adapter.h" @@ -26,8 +27,10 @@ CNodePtr MatMulFusionInoutTest::AddMatMul(const FuncGraphPtr &graph, const AnfNo const ActivationType &act_type, const std::string &name) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create MatMul primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(false, false, act_type); - auto matmul_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto matmul_primitive = NewValueNode(prim_c); auto matmul = graph->NewCNode({matmul_primitive, input1, input2}); MS_CHECK_TRUE_MSG(matmul != nullptr, nullptr, "create MatMul failed"); matmul->set_fullname_with_scope(name); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_mul_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_mul_fusion_inout_test.cc index 53a4a16f4e..af348652dd 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_mul_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/matmul_mul_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "tools/optimizer/fusion/matmul_mul_fusion.h" #include "test/ut/tools/optimizer/fusion/fusion_inout_test/conv_fusion_inout_test.h" @@ -55,8 +56,10 @@ class MatmulMulFusionInoutTest : public FusionInoutTest { AnfNodePtr param = AddParameter(graph_, 0, {in_}, kNumberTypeFloat32, name); auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create MulFusion primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(ActivationType::NO_ACTIVATION); - auto add_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto add_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(add_primitive != nullptr, nullptr); auto add_fusion = graph->NewCNode({add_primitive, input, param}); MS_CHECK_TRUE_MSG(add_fusion != nullptr, nullptr, "create AddFusion failed"); @@ -69,8 +72,10 @@ class MatmulMulFusionInoutTest : public FusionInoutTest { AnfNodePtr bias = AddParameter(graph_, 0, {in_}, kNumberTypeFloat32, "graph_" + name + "_input3"); auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create MatMul primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(false, false, ActivationType::NO_ACTIVATION); - auto add_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto add_primitive = NewValueNode(prim_c); MS_CHECK_TRUE_RET(add_primitive != nullptr, nullptr); auto add_fusion = graph->NewCNode({add_primitive, input1, input2, bias}); MS_CHECK_TRUE_MSG(add_fusion != nullptr, nullptr, "create AddFusion failed"); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/trans_matmul_fusion_inout_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/trans_matmul_fusion_inout_test.cc index f85ccdd821..6344bf08cb 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/trans_matmul_fusion_inout_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/fusion_inout_test/trans_matmul_fusion_inout_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "tools/optimizer/fusion/transpose_matmul_fusion.h" #include "tools/optimizer/common/gllo_utils.h" @@ -76,8 +77,10 @@ class TransMatMulFusionInoutTest : public MatMulFusionInoutTest { const std::string &name) { auto prim = std::make_unique(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "create Act primitivec failed"); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_MSG(prim_c != nullptr, nullptr, "prim_c is nullptr"); prim->Init(); - auto trans_primitive = NewValueNode(std::shared_ptr(prim.release())); + auto trans_primitive = NewValueNode(prim_c); auto perm = opt::BuildIntVecParameterNode(graph, perm_val, name + "_perm"); auto transpose = graph->NewCNode({trans_primitive, perm}); MS_CHECK_TRUE_MSG(transpose != nullptr, nullptr, "create Transpose failed"); diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/matmul_mul_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/matmul_mul_fusion_test.cc index 2074ba6824..b1b2a2aa3d 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/matmul_mul_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/matmul_mul_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include "schema/inner/model_generated.h" #include "include/model.h" diff --git a/mindspore/lite/test/ut/tools/optimizer/fusion/trans_matmul_fusion_test.cc b/mindspore/lite/test/ut/tools/optimizer/fusion/trans_matmul_fusion_test.cc index 116769d938..7cb5ada7c7 100644 --- a/mindspore/lite/test/ut/tools/optimizer/fusion/trans_matmul_fusion_test.cc +++ b/mindspore/lite/test/ut/tools/optimizer/fusion/trans_matmul_fusion_test.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include #include "schema/inner/model_generated.h" diff --git a/mindspore/lite/tools/anf_exporter/anf_exporter.cc b/mindspore/lite/tools/anf_exporter/anf_exporter.cc index d7da0bb74f..c2cea8e6a7 100644 --- a/mindspore/lite/tools/anf_exporter/anf_exporter.cc +++ b/mindspore/lite/tools/anf_exporter/anf_exporter.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/anf_exporter/anf_exporter.h" #include #include @@ -24,6 +25,7 @@ #include "tools/converter/converter_flags.h" #include "abstract/abstract_value.h" #include "mindspore/core/ir/primitive.h" +#include "mindspore/core/ops/op_name.h" #include "mindspore/core/ops/op_utils.h" #include "ops/fusion/partial_fusion.h" #include "ops/call.h" @@ -113,7 +115,7 @@ int AnfExporter::SetPostTrainOutputTensorType(const std::unique_ptrname() != mindspore::ops::kNameQuantDTypeCast) { first_tensor_output->dataType = kNumberTypeInt8; } else { - auto primc = primitive->cast>(); + auto primc = api::MakeShared(primitive); MS_CHECK_TRUE_MSG(primc != nullptr, RET_ERROR, "cast ptr failed"); if (primc->get_dst_t() != kNumberTypeFloat32) { first_tensor_output->dataType = kNumberTypeInt8; @@ -639,7 +641,11 @@ schema::MetaGraphT *AnfExporter::Export(const FuncGraphPtr &func_graph, bool kee MS_CHECK_TRUE_MSG(meta_graphT != nullptr, nullptr, "meta_graphT is nullptr"); auto fmk = func_graph->get_attr("fmk"); MS_CHECK_TRUE_MSG(fmk != nullptr, nullptr, "fmk is nullptr"); - meta_graphT->fmkType = GetValue(fmk); + if (fmk->isa()) { + meta_graphT->fmkType = GetValue(fmk); + } else { + meta_graphT->fmkType = GetValue(fmk); + } graph_inputs_ = func_graph->get_inputs(); diff --git a/mindspore/lite/tools/anf_exporter/fetch_content.cc b/mindspore/lite/tools/anf_exporter/fetch_content.cc index b7c328a2a3..dc9968a083 100644 --- a/mindspore/lite/tools/anf_exporter/fetch_content.cc +++ b/mindspore/lite/tools/anf_exporter/fetch_content.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/anf_exporter/fetch_content.h" #include #include @@ -30,6 +31,8 @@ #include "tools/common/node_util.h" #include "src/ops/ops_utils.h" #include "src/ops/populate/populate_register.h" +#include "mindapi/base/format.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/common/func_graph_subgraph.cc b/mindspore/lite/tools/common/func_graph_subgraph.cc index 6bca0c164a..787a3bd651 100644 --- a/mindspore/lite/tools/common/func_graph_subgraph.cc +++ b/mindspore/lite/tools/common/func_graph_subgraph.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/common/func_graph_subgraph.h" #include #include @@ -445,8 +446,10 @@ int SubGraph::CreatePartialInBelongAnf() { auto graph_value_node = NewValueNode(func_graph); MS_CHECK_TRUE_MSG(partial_prim != nullptr, RET_NULL_PTR, "partial_prim is nullptr"); MS_CHECK_TRUE_MSG(graph_value_node != nullptr, RET_NULL_PTR, "graph_value_node is nullptr"); + auto partial_prim_c = partial_prim->GetPrim(); + MS_CHECK_TRUE_MSG(partial_prim_c != nullptr, RET_NULL_PTR, "partial_prim_c is nullptr"); partial_inputs.insert(partial_inputs.begin(), graph_value_node); - auto partial_cnode = belong_anf_->NewCNode(partial_prim, partial_inputs); + auto partial_cnode = belong_anf_->NewCNode(partial_prim_c, partial_inputs); MS_CHECK_TRUE_MSG(partial_cnode != nullptr, RET_NULL_PTR, "partial_cnode is nullptr"); partial_cnode->set_fullname_with_scope(graph_name + "/partial"); for (size_t i = 0; i < partial_inputs.size(); ++i) { diff --git a/mindspore/lite/tools/common/graph_util.cc b/mindspore/lite/tools/common/graph_util.cc index d85233d664..939779e6c8 100644 --- a/mindspore/lite/tools/common/graph_util.cc +++ b/mindspore/lite/tools/common/graph_util.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/common/graph_util.h" #include #include @@ -29,6 +30,7 @@ #include "src/common/utils.h" #include "nnacl/op_base.h" #include "ops/make_tuple.h" +#include "tools/converter/converter_context.h" namespace mindspore { namespace lite { @@ -49,8 +51,10 @@ int SetFuncGraphOutput(const FuncGraphPtr &graph, const std::vector MS_LOG(DEBUG) << "new MakeTuple failed"; return lite::RET_NULL_PTR; } - auto make_tuple_cnode = graph->NewCNode(make_tuple_prim_ptr, outputs); - if (make_tuple_prim_ptr == nullptr) { + auto make_tuple_prim_c = make_tuple_prim_ptr->GetPrim(); + MS_CHECK_TRUE_MSG(make_tuple_prim_c != nullptr, lite::RET_NULL_PTR, "make_tuple_prim_c is nullptr"); + auto make_tuple_cnode = graph->NewCNode(make_tuple_prim_c, outputs); + if (make_tuple_cnode == nullptr) { MS_LOG(DEBUG) << "new cnode failed"; return lite::RET_NULL_PTR; } diff --git a/mindspore/lite/tools/common/node_util.cc b/mindspore/lite/tools/common/node_util.cc index 6ef00b7cb1..d1bee1b0c2 100644 --- a/mindspore/lite/tools/common/node_util.cc +++ b/mindspore/lite/tools/common/node_util.cc @@ -389,7 +389,9 @@ bool IsMakeTuple(const AnfNodePtr &node) { ValueNodePtr GetPartialFusionPrim() { auto partial_prim = std::make_shared(); MS_CHECK_TRUE_MSG(partial_prim != nullptr, nullptr, "partial_prim is nullptr"); - ValueNodePtr partial_anf_prim = NewValueNode(partial_prim); + auto partial_prim_c = partial_prim->GetPrim(); + MS_CHECK_TRUE_MSG(partial_prim_c != nullptr, nullptr, "partial_prim_c is nullptr"); + ValueNodePtr partial_anf_prim = NewValueNode(partial_prim_c); MS_CHECK_TRUE_MSG(partial_anf_prim != nullptr, nullptr, "partial_anf_prim is nullptr"); return partial_anf_prim; } @@ -397,7 +399,9 @@ ValueNodePtr GetPartialFusionPrim() { ValueNodePtr GetSwitchAnfPrim() { auto switch_prim = std::make_shared(); MS_CHECK_TRUE_MSG(switch_prim != nullptr, nullptr, "switch_prim is nullptr"); - ValueNodePtr switch_anf_prim = NewValueNode(switch_prim); + auto switch_prim_c = switch_prim->GetPrim(); + MS_CHECK_TRUE_MSG(switch_prim_c != nullptr, nullptr, "switch_prim_c is nullptr"); + ValueNodePtr switch_anf_prim = NewValueNode(switch_prim_c); MS_CHECK_TRUE_MSG(switch_prim != nullptr, nullptr, "switch_prim is nullptr"); return switch_anf_prim; } @@ -405,7 +409,9 @@ ValueNodePtr GetSwitchAnfPrim() { ValueNodePtr GetCallAnfPrim() { auto call_prim = std::make_shared(); MS_CHECK_TRUE_MSG(call_prim != nullptr, nullptr, "call_prim is nullptr"); - ValueNodePtr call_anf_prim = NewValueNode(call_prim); + auto call_prim_c = call_prim->GetPrim(); + MS_CHECK_TRUE_MSG(call_prim_c != nullptr, nullptr, "call_prim_c is nullptr"); + ValueNodePtr call_anf_prim = NewValueNode(call_prim_c); MS_CHECK_TRUE_MSG(call_anf_prim != nullptr, nullptr, "call_anf_prim is nullptr"); return call_anf_prim; } diff --git a/mindspore/lite/tools/common/node_util.h b/mindspore/lite/tools/common/node_util.h index 722244d536..2bb826372f 100644 --- a/mindspore/lite/tools/common/node_util.h +++ b/mindspore/lite/tools/common/node_util.h @@ -27,6 +27,7 @@ #include "src/tensor.h" #include "include/errorcode.h" #include "securec/include/securec.h" +#include "ops/primitive_c.h" #include "tools/optimizer/common/gllo_utils.h" namespace mindspore { diff --git a/mindspore/lite/tools/converter/adapter/acl/acl_pass.h b/mindspore/lite/tools/converter/adapter/acl/acl_pass.h index 9a92cb5d10..35ecd091d9 100644 --- a/mindspore/lite/tools/converter/adapter/acl/acl_pass.h +++ b/mindspore/lite/tools/converter/adapter/acl/acl_pass.h @@ -17,6 +17,7 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_ADAPTER_ACL_PASS_H #define MINDSPORE_LITE_TOOLS_CONVERTER_ADAPTER_ACL_PASS_H +#define USE_DEPRECATED_API #include #include "backend/common/optimizer/pass.h" #include "tools/converter/converter_flags.h" diff --git a/mindspore/lite/tools/converter/adapter/acl/common/utils.h b/mindspore/lite/tools/converter/adapter/acl/common/utils.h index c20fad63cf..6eb487ef67 100644 --- a/mindspore/lite/tools/converter/adapter/acl/common/utils.h +++ b/mindspore/lite/tools/converter/adapter/acl/common/utils.h @@ -18,13 +18,17 @@ #define TOOLS_CONVERTER_ADAPTER_ACL_COMMON_COMMON_UTILS_H #include +#include #include #include "include/errorcode.h" #include "ir/anf.h" #include "ir/dtype/type_id.h" +#include "ops/base_operator.h" namespace mindspore { namespace lite { +using BaseOperatorPtr = std::shared_ptr; + namespace acl { STATUS GetShapeVectorFromCNode(const mindspore::CNodePtr &cnode, std::vector *shape_vector); diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/activation_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/activation_mapper.cc index a14f2b2801..b6c769ae9d 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/activation_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/activation_mapper.cc @@ -14,10 +14,12 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/activation_mapper.h" #include #include #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" +#include "tools/converter/adapter/acl/common/utils.h" #include "ops/elu.h" #include "ops/gelu.h" #include "ops/leaky_relu.h" @@ -26,11 +28,13 @@ #include "ops/sigmoid.h" #include "ops/tanh.h" #include "nnacl/op_base.h" +#include "src/common/log_util.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { STATUS ActivationMapper::Mapper(const CNodePtr &cnode) { - static std::map activation_type_map = { + static std::map activation_type_map = { {mindspore::ELU, std::make_shared()}, {mindspore::GELU, std::make_shared()}, {mindspore::RELU, std::make_shared()}, @@ -44,17 +48,19 @@ STATUS ActivationMapper::Mapper(const CNodePtr &cnode) { MS_LOG(ERROR) << "Get primitive from cnode failed."; return lite::RET_ERROR; } - auto activate_prim = dynamic_cast(src_prim.get()); + auto activate_prim = mindspore::api::MakeShared(src_prim); MS_CHECK_TRUE_MSG(activate_prim != nullptr, lite::RET_ERROR, "Dynamic cast activation failed."); PrimitivePtr dst_prim = nullptr; ActivationType type = activate_prim->get_activation_type(); if (activation_type_map.find(type) != activation_type_map.end()) { - dst_prim = activation_type_map[type]; + auto dest_op = activation_type_map[type]; + MS_CHECK_TRUE_MSG(dest_op != nullptr, lite::RET_ERROR, "Activation op failed."); + dst_prim = dest_op->GetPrim(); } else { MS_LOG(ERROR) << "Type " << static_cast(type) << " is unsupported."; return lite::RET_ERROR; } - MS_ASSERT(dst_prim != nullptr); + MS_CHECK_TRUE_MSG(dst_prim != nullptr, lite::RET_ERROR, "Dst prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); value_node->set_value(dst_prim); return lite::RET_OK; diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/argmax_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/argmax_fusion_mapper.cc index 65020757be..170716d95c 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/argmax_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/argmax_fusion_mapper.cc @@ -19,6 +19,7 @@ #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "src/common/log_util.h" #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/arithmetic_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/arithmetic_mapper.cc index 05a9c30113..b299a967de 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/arithmetic_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/arithmetic_mapper.cc @@ -14,21 +14,26 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/arithmetic_mapper.h" #include #include #include #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" +#include "tools/converter/adapter/acl/common/utils.h" #include "src/common/log_util.h" #include "ops/real_div.h" +#include "ops/op_utils.h" +#include "nnacl/op_base.h" namespace mindspore { namespace lite { -static const std::map kDivTypeMap = {{"Div", std::make_shared()}, - {"RealDiv", std::make_shared()}}; +static const std::map kDivTypeMap = {{"Div", std::make_shared()}, + {"RealDiv", std::make_shared()}}; STATUS AddFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::Add op_add; + auto dst_prim = op_add.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "AddFusion mapper failed."; return RET_ERROR; @@ -51,7 +56,9 @@ STATUS DivFusionMapper::Mapper(const CNodePtr &cnode) { } PrimitivePtr dst_prim = nullptr; if (kDivTypeMap.find(original_name) != kDivTypeMap.end()) { - dst_prim = kDivTypeMap.at(original_name); + auto dst_op = kDivTypeMap.at(original_name); + MS_CHECK_TRUE_MSG(dst_op != nullptr, lite::RET_ERROR, "Div op is nullptr."); + dst_prim = dst_op->GetPrim(); } CHECK_NULL_RETURN(dst_prim); dst_prim->SetAttrs(src_prim->attrs()); @@ -60,7 +67,8 @@ STATUS DivFusionMapper::Mapper(const CNodePtr &cnode) { } STATUS MulFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::Mul mul_op; + auto dst_prim = mul_op.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "MulFusion mapper failed."; return RET_ERROR; @@ -69,7 +77,8 @@ STATUS MulFusionMapper::Mapper(const CNodePtr &cnode) { } STATUS PowFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::Pow pow_op; + auto dst_prim = pow_op.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "PowFusion mapper failed."; return RET_ERROR; @@ -78,7 +87,8 @@ STATUS PowFusionMapper::Mapper(const CNodePtr &cnode) { } STATUS SubFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::Sub sub_op; + auto dst_prim = sub_op.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "SubFusion mapper failed."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/avgpool_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/avgpool_fusion_mapper.cc index 6ad4449471..ffe819f7ed 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/avgpool_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/avgpool_fusion_mapper.cc @@ -20,6 +20,8 @@ #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" #include "include/registry/converter_context.h" #include "src/common/log_util.h" +#include "ops/op_utils.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { @@ -46,6 +48,11 @@ STATUS AvgPoolFusionMapper::Mapper(const CNodePtr &cnode) { } void AvgPoolFusionMapper::CreateTargetPrim(const PrimitivePtr &src_prim, PrimitivePtr *dst_prim, int fmk_type) { + if (dst_prim == nullptr) { + MS_LOG(ERROR) << "Target prim is nullptr."; + return; + } + if (fmk_type == converter::kFmkTypeCaffe) { *dst_prim = std::make_shared(); } else if (fmk_type == converter::kFmkTypeOnnx) { @@ -56,7 +63,8 @@ void AvgPoolFusionMapper::CreateTargetPrim(const PrimitivePtr &src_prim, Primiti *dst_prim = std::make_shared(); } } else { - *dst_prim = std::make_shared(); + ops::AvgPool dst_node; + *dst_prim = dst_node.GetPrim(); } } diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/batchnorm_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/batchnorm_mapper.cc index 6483b6afbf..eec1e2406b 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/batchnorm_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/batchnorm_mapper.cc @@ -19,6 +19,7 @@ #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" #include "include/registry/converter_context.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/cast_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/cast_mapper.cc index 1d49a837f8..7433010ed9 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/cast_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/cast_mapper.cc @@ -18,6 +18,7 @@ #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "tools/converter/adapter/acl/common/utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_backprop_input_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_backprop_input_mapper.cc index 0b34e4d502..6641e8d6f0 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_backprop_input_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_backprop_input_mapper.cc @@ -20,6 +20,7 @@ #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "src/common/log_util.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_fusion_mapper.cc index 9730dcda92..c9a0efa4f0 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_fusion_mapper.cc @@ -19,6 +19,7 @@ #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" #include "src/common/log_util.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { @@ -36,7 +37,8 @@ STATUS Conv2DFusionMapper::Mapper(const CNodePtr &cnode) { } PrimitivePtr dst_prim = nullptr; if (!is_depth_wise) { - dst_prim = std::make_shared(); + ops::Conv2D conv2d_op; + dst_prim = conv2d_op.GetPrim(); } else { dst_prim = std::make_shared(); } diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_transpose_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_transpose_fusion_mapper.cc index db3b39c424..ae993b41a1 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_transpose_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/conv2d_transpose_fusion_mapper.cc @@ -22,6 +22,7 @@ #include "include/registry/converter_context.h" #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/conv_base_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/conv_base_mapper.cc index 7892f72cac..52fbee7431 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/conv_base_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/conv_base_mapper.cc @@ -19,6 +19,7 @@ #include #include "ops/op_utils.h" #include "nnacl/op_base.h" +#include "utils/check_convert_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/gather_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/gather_fusion_mapper.cc index 9e01aa330b..929358aff3 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/gather_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/gather_fusion_mapper.cc @@ -14,16 +14,18 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/gather_fusion_mapper.h" #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "tools/converter/adapter/acl/common/utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { namespace { constexpr size_t kNameGatherInputNum = 4; -} +} // namespace STATUS GatherMapper::Mapper(const CNodePtr &cnode) { MS_CHECK_TRUE_MSG(cnode != nullptr, lite::RET_ERROR, "Cnode is nullptr."); diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/matmul_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/matmul_fusion_mapper.cc index 6717502d07..9d52d56fcc 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/matmul_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/matmul_fusion_mapper.cc @@ -17,11 +17,13 @@ #include "tools/converter/adapter/acl/mapper/matmul_fusion_mapper.h" #include #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { STATUS MatMulFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::MatMul mat_mul; + auto dst_prim = mat_mul.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "MatMulFusion mapper failed."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/maxpool_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/maxpool_fusion_mapper.cc index 9b1d5ba8d7..4318d3813c 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/maxpool_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/maxpool_fusion_mapper.cc @@ -20,6 +20,7 @@ #include "tools/converter/adapter/acl/mapper/tbe_op_def.h" #include "include/registry/converter_context.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { @@ -39,7 +40,8 @@ STATUS MaxPoolFusionMapper::Mapper(const CNodePtr &cnode) { } else if (fmk_type == converter::kFmkTypeOnnx) { dst_prim = std::make_shared(); } else { - dst_prim = std::make_shared(); + ops::MaxPool max_pool_op; + dst_prim = max_pool_op.GetPrim(); } MS_CHECK_TRUE_MSG(dst_prim != nullptr, lite::RET_ERROR, "Get primitive by fmk type failed."); dst_prim->SetAttrs(src_prim->attrs()); diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/prelu_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/prelu_fusion_mapper.cc index 72d0b030f5..77d36d34c0 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/prelu_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/prelu_fusion_mapper.cc @@ -14,14 +14,17 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/prelu_fusion_mapper.h" #include #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { STATUS PReluFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::PReLU prelu_op; + auto dst_prim = prelu_op.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "PReluFusion mapper failed."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/reduce_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/reduce_fusion_mapper.cc index c267891cba..3baa5d060f 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/reduce_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/reduce_fusion_mapper.cc @@ -44,9 +44,11 @@ STATUS ReduceFusionMapper::Mapper(const CNodePtr &cnode) { int64_t mode = GetValue(attr_val); PrimitivePtr dst_prim = nullptr; if (mode == static_cast(ReduceMode::Reduce_Sum)) { - dst_prim = std::make_shared(); + ops::ReduceSum reduce_sum_op; + dst_prim = reduce_sum_op.GetPrim(); } else if (mode == static_cast(ReduceMode::Reduce_Mean)) { - dst_prim = std::make_shared(); + ops::ReduceMean reduce_mean_op; + dst_prim = reduce_mean_op.GetPrim(); } CHECK_NULL_RETURN(dst_prim); dst_prim->SetAttrs(src_prim->attrs()); diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/resize_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/resize_mapper.cc index 9c2b631567..c7d1e962c4 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/resize_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/resize_mapper.cc @@ -25,7 +25,7 @@ namespace mindspore { namespace lite { namespace { constexpr auto kNameInputNum = 3; -} +} // namespace STATUS ResizeMapper::Mapper(const CNodePtr &cnode) { if (cnode->inputs().size() != kNameInputNum) { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/scale_fusion_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/scale_fusion_mapper.cc index 8444e2093a..d1ab2ba3b8 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/scale_fusion_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/scale_fusion_mapper.cc @@ -14,14 +14,17 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/scale_fusion_mapper.h" #include #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { STATUS ScaleFusionMapper::Mapper(const CNodePtr &cnode) { - auto dst_prim = std::make_shared(); + ops::Scale scale_op; + auto dst_prim = scale_op.GetPrim(); if (MoveAttrMap(cnode, dst_prim) != RET_OK) { MS_LOG(ERROR) << "ScaleFusion mapper failed."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/spatial_node_adapter.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/spatial_node_adapter.cc index 4702f105f3..8af63abe7d 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/spatial_node_adapter.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/spatial_node_adapter.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/spatial_node_adapter.h" #include #include @@ -46,7 +47,8 @@ CNodePtr CreateTupleGetItemNode(const FuncGraphPtr &func_graph, const CNodePtr & CNodePtr get_item_cnode = nullptr; auto tuple_get_item_prim_ptr = std::make_shared(); MS_CHECK_TRUE_MSG(tuple_get_item_prim_ptr != nullptr, nullptr, "New TupleGetItem failed"); - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); + auto tuple_get_item_prim_c = tuple_get_item_prim_ptr->GetPrim(); + auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_c); MS_CHECK_TRUE_MSG(tuple_get_item_prim != nullptr, nullptr, "tuple_prim is nullptr."); auto get_item_value = NewValueNode(MakeValue(0)); MS_CHECK_TRUE_MSG(get_item_value != nullptr, nullptr, "item_value is nullptr."); diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/stack_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/stack_mapper.cc index d40ceebdc9..77f2f54d8c 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/stack_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/stack_mapper.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { namespace { constexpr auto kNameNum = "num"; -} +} // namespace STATUS StackMapper::Mapper(const CNodePtr &cnode) { if (AddAttrForDynInputPrimitive(cnode, kNameNum) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/stridedslice_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/stridedslice_mapper.cc index 301d6ebcf8..bb918782ad 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/stridedslice_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/stridedslice_mapper.cc @@ -33,8 +33,8 @@ STATUS StridedSliceMapper::Mapper(const CNodePtr &cnode) { } auto attr_val = src_prim->GetAttr(ops::kFmkType); - int fmk_type = attr_val != nullptr ? GetValue(attr_val) : converter::kFmkTypeTf; - if (fmk_type == converter::kFmkTypeOnnx) { + int64_t fmk_type = attr_val != nullptr ? GetValue(attr_val) : converter::kFmkTypeTf; + if (static_cast(fmk_type) == converter::kFmkTypeOnnx) { auto dst_prim = std::make_shared(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, lite::RET_ERROR, "dst_prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); diff --git a/mindspore/lite/tools/converter/adapter/acl/mapper/transpose_mapper.cc b/mindspore/lite/tools/converter/adapter/acl/mapper/transpose_mapper.cc index 7572c56520..338d47e27f 100644 --- a/mindspore/lite/tools/converter/adapter/acl/mapper/transpose_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/acl/mapper/transpose_mapper.cc @@ -14,18 +14,20 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/mapper/transpose_mapper.h" #include #include #include "tools/converter/adapter/acl/mapper/primitive_mapper_register.h" #include "tools/converter/adapter/acl/common/utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { namespace { constexpr size_t kCommonInputNum = 3; -} +} // namespace STATUS TransposeMapper::Mapper(const CNodePtr &cnode) { MS_CHECK_TRUE_MSG(cnode != nullptr, lite::RET_ERROR, "cnode is nullptr."); diff --git a/mindspore/lite/tools/converter/adapter/acl/src/acl_pass_impl.cc b/mindspore/lite/tools/converter/adapter/acl/src/acl_pass_impl.cc index c5f500d66c..29075e0822 100644 --- a/mindspore/lite/tools/converter/adapter/acl/src/acl_pass_impl.cc +++ b/mindspore/lite/tools/converter/adapter/acl/src/acl_pass_impl.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/adapter/acl/src/acl_pass_impl.h" #include #include @@ -28,6 +29,7 @@ #include "tools/converter/adapter/acl/src/acl_model_process.h" #include "include/registry/pass_registry.h" #include "ops/custom.h" +#include "ops/op_utils.h" #include "ops/tuple_get_item.h" #include "base/core_ops.h" #include "cxx_api/model/acl/model_converter.h" @@ -537,8 +539,9 @@ CNodePtr AclPassImpl::CreateCustomNode(const FuncGraphPtr &func_graph) { auto prim = std::make_shared(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "New custom op failed."); prim->set_type(kCustomPrimTypeACL); + auto prim_c = prim->GetPrim(); auto graph_input = func_graph->get_inputs(); - CNodePtr custom_node = func_graph->NewCNode(prim, graph_input); + CNodePtr custom_node = func_graph->NewCNode(prim_c, graph_input); MS_CHECK_TRUE_MSG(custom_node != nullptr, nullptr, "Custom cnode failed."); custom_node->set_fullname_with_scope(kCustomNodeName); custom_node->add_input(om_parameter_); @@ -580,7 +583,8 @@ STATUS AclPassImpl::ModifyGraphByCustomNode(const FuncGraphPtr &func_graph, cons MS_LOG(ERROR) << "New TupleGetItem failed for output " << j; return lite::RET_ERROR; } - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); + auto tuple_get_item_prim_ptr_c = tuple_get_item_prim_ptr->GetPrim(); + auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr_c); MS_CHECK_TRUE_MSG(tuple_get_item_prim != nullptr, lite::RET_ERROR, "item_prim is nullptr."); auto get_item_value = NewValueNode(MakeValue(j)); MS_CHECK_TRUE_MSG(get_item_value != nullptr, lite::RET_ERROR, "item_value is nullptr."); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.cc index 7f52284459..affa9c2697 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.cc @@ -16,8 +16,8 @@ #include "checker/activation_checker.h" #include -#include #include +#include "mindapi/base/types.h" #include "common/op_enum.h" namespace mindspore { @@ -27,12 +27,12 @@ const std::unordered_set kSupportedActivationTypes = { ActivationType::RELU, ActivationType::RELU6, ActivationType::LEAKY_RELU, ActivationType::SIGMOID, ActivationType::TANH, ActivationType::HSWISH, ActivationType::HARD_TANH, ActivationType::ELU}; } // namespace -bool ActivationChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ActivationChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(INFO) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; @@ -42,7 +42,7 @@ bool ActivationChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format MS_LOG(ERROR) << "kActivationType attr is nullptr."; return false; } - auto activation_type = static_cast(GetValue(value_ptr)); + auto activation_type = static_cast(api::GetValue(value_ptr)); if (kSupportedActivationTypes.find(activation_type) == kSupportedActivationTypes.end()) { MS_LOG(WARNING) << "Not supported activation type: " << activation_type << ", will turn it to custom op. " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.h index 3c9e3cbf30..a70dff265d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/activation_checker.h @@ -26,7 +26,7 @@ class ActivationChecker : public OpChecker { public: ActivationChecker() : OpChecker("Activation") {} ~ActivationChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.cc index cdfea97964..f4b573fcb1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.cc @@ -23,18 +23,18 @@ namespace mindspore { namespace dpico { -bool ArgMaxChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ArgMaxChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr." << op->fullname_with_scope(); return false; } if (primitive->GetAttr(ops::kTopK) != nullptr) { - auto top_k = static_cast(GetValue(primitive->GetAttr(ops::kTopK))); + auto top_k = static_cast(api::GetValue(primitive->GetAttr(ops::kTopK))); if (top_k != 1) { MS_LOG(WARNING) << "top_k value only supports 1 for dpico. " << op->fullname_with_scope(); return false; @@ -42,7 +42,7 @@ bool ArgMaxChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for } if (primitive->GetAttr(ops::kAxis) != nullptr) { - auto axis = GetValue(primitive->GetAttr(ops::kAxis)); + auto axis = api::GetValue(primitive->GetAttr(ops::kAxis)); if (axis <= kAxisLowerBound || axis > kAxisUpperBound || axis == 0) { MS_LOG(WARNING) << op->fullname_with_scope() << "'s axis should in range (-4, 0) and (0, 3], but in fact it's " << axis; @@ -50,7 +50,7 @@ bool ArgMaxChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for } } if (primitive->GetAttr(dpico::kSelectLastIndex) != nullptr) { - auto select_last_index = GetValue(primitive->GetAttr(dpico::kSelectLastIndex)); + auto select_last_index = api::GetValue(primitive->GetAttr(dpico::kSelectLastIndex)); if (!select_last_index) { MS_LOG(WARNING) << "select_last_index value only supports true for dpico. " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.h index 36575b7ff4..37b54f050d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/argmax_checker.h @@ -26,7 +26,7 @@ class ArgMaxChecker : public OpChecker { public: ArgMaxChecker() : OpChecker("ArgMax") {} ~ArgMaxChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.cc index 8dbbad60d0..14dc72fe9d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.cc @@ -14,22 +14,22 @@ * limitations under the License. */ +#include "checker/arithmetic_checker.h" #include #include #include #include "common/anf_util.h" #include "common/check_base.h" #include "common/op_enum.h" -#include "checker/arithmetic_checker.h" namespace mindspore { namespace dpico { -bool ArithmeticChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ArithmeticChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto prim = GetValueNode(op->input(0)); + auto prim = api::GetValueNode(op->input(0)); MS_CHECK_TRUE_MSG(prim != nullptr, false, "prim is nullptr." << op->fullname_with_scope()); if (op->inputs().size() != kDims3) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.h index dba9a5d9e7..bac4381ea4 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/arithmetic_checker.h @@ -26,7 +26,7 @@ class ArithmeticChecker : public OpChecker { public: ArithmeticChecker() : OpChecker("Arithmetic") {} ~ArithmeticChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.cc index 94ef5de3dc..3fd49e1219 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.cc @@ -22,18 +22,18 @@ namespace mindspore { namespace dpico { -bool BatchNormChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool BatchNormChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive_c is nullptr." << op->fullname_with_scope(); return false; } if (primitive->GetAttr(kUseGlobalStats) != nullptr) { - auto use_global_stats = GetValue(primitive->GetAttr(kUseGlobalStats)); + auto use_global_stats = api::GetValue(primitive->GetAttr(kUseGlobalStats)); if (!use_global_stats) { MS_LOG(WARNING) << "use global stats attr is false, which is not supported by dpico. " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.h index 070cb2183d..3fbe863142 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/batchnorm_checker.h @@ -26,7 +26,7 @@ class BatchNormChecker : public OpChecker { public: BatchNormChecker() : OpChecker("BatchNorm") {} ~BatchNormChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.cc index dce93707d9..fe540ac2fb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace dpico { -bool CommonChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool CommonChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.h index 1bf3298f81..e719471fa0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/common_checker.h @@ -26,7 +26,7 @@ class CommonChecker : public OpChecker { public: CommonChecker() : OpChecker("Common") {} ~CommonChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.cc index 9db529b7e4..741b992769 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.cc @@ -17,6 +17,7 @@ #include "checker/concat_checker.h" #include #include +#include #include "common/anf_util.h" #include "common/op_enum.h" @@ -40,8 +41,8 @@ bool CheckConcatInputW(const ShapeVector &input_shape, int64_t axis, int64_t inp return supported; } } // namespace -bool ConcatChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool ConcatChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr." << op->fullname_with_scope(); return false; @@ -49,7 +50,7 @@ bool ConcatChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for int64_t axis = 0; auto axis_ptr = primitive->GetAttr(ops::kAxis); if (axis_ptr != nullptr) { - axis = GetValue(axis_ptr); + axis = api::GetValue(axis_ptr); if (axis < kAxisLowerBound || axis > kAxisUpperBound) { MS_LOG(WARNING) << op->fullname_with_scope() << "'s axis should in range [-4, 3], but in fact it's " << axis; return false; @@ -65,7 +66,7 @@ bool ConcatChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for std::vector input_shape; for (size_t i = 1; i < op->inputs().size(); i++) { auto input = op->input(i); - if (input->cast() != nullptr && input->cast()->has_default()) { + if (input->cast() != nullptr && input->cast()->has_default()) { MS_LOG(WARNING) << "there is offline data in concat, which dpico is unsupported. " << op->fullname_with_scope(); return false; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.h index 7259e5c87c..5a759283fc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/concat_checker.h @@ -26,7 +26,7 @@ class ConcatChecker : public OpChecker { public: ConcatChecker() : OpChecker("Concat") {} ~ConcatChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.cc index 801379cb7e..6e63dc760b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.cc @@ -20,6 +20,8 @@ #include "common/anf_util.h" #include "common/op_enum.h" #include "common/check_base.h" +#include "ops/fusion/conv2d_fusion.h" +#include "ops/fusion/conv2d_transpose_fusion.h" namespace mindspore { namespace dpico { @@ -33,10 +35,10 @@ constexpr int kMaxDilationAndKernelProd = 15; constexpr int kConvUpperBound = 2048; constexpr float kCoefficient = 16.0; -bool CheckOutChannel(const PrimitivePtr &primitive) { +bool CheckOutChannel(const api::PrimitivePtr &primitive) { auto out_channel_ptr = primitive->GetAttr(ops::kOutChannel); if (out_channel_ptr != nullptr) { - auto out_channel_data = GetValue(out_channel_ptr); + auto out_channel_data = api::GetValue(out_channel_ptr); if (out_channel_data < 1 || out_channel_data > kMaxNumOutput) { MS_LOG(WARNING) << "out_channel:" << out_channel_data << " is unsupported by dpico."; return false; @@ -44,10 +46,10 @@ bool CheckOutChannel(const PrimitivePtr &primitive) { } return true; } -bool CheckGroup(const PrimitivePtr &primitive) { +bool CheckGroup(const api::PrimitivePtr &primitive) { auto group_ptr = primitive->GetAttr(ops::kGroup); if (group_ptr != nullptr) { - auto group_data = GetValue(group_ptr); + auto group_data = api::GetValue(group_ptr); if (group_data < 1 || group_data > kMaxGroupNum) { MS_LOG(WARNING) << "group:" << group_data << " is unsupported by dpico."; return false; @@ -55,10 +57,10 @@ bool CheckGroup(const PrimitivePtr &primitive) { } return true; } -bool CheckPadList(const PrimitivePtr &primitive) { +bool CheckPadList(const api::PrimitivePtr &primitive) { auto pad_ptr = primitive->GetAttr(ops::kPadList); if (pad_ptr != nullptr) { - auto pad_data = GetValue>(pad_ptr); + auto pad_data = api::GetValue>(pad_ptr); if (pad_data.size() > kDims3) { if (pad_data[0] < 0 || pad_data[0] > kMaxPadSize || pad_data[1] < 0 || pad_data[1] > kMaxPadSize || pad_data[kAxis2] < 0 || pad_data[kAxis2] > kMaxPadSize || pad_data[kAxis3] < 0 || @@ -70,10 +72,10 @@ bool CheckPadList(const PrimitivePtr &primitive) { } return true; } -bool CheckKernelSize(const PrimitivePtr &primitive) { +bool CheckKernelSize(const api::PrimitivePtr &primitive) { auto kernel_ptr = primitive->GetAttr(ops::kKernelSize); if (kernel_ptr != nullptr) { - auto kernel_data = GetValue>(kernel_ptr); + auto kernel_data = api::GetValue>(kernel_ptr); if (kernel_data.size() > 1) { if (kernel_data[0] < 1 || kernel_data[0] > kMaxKernelSize || kernel_data[1] < 1 || kernel_data[1] > kMaxKernelSize) { @@ -84,10 +86,10 @@ bool CheckKernelSize(const PrimitivePtr &primitive) { } return true; } -bool CheckStride(const PrimitivePtr &primitive) { +bool CheckStride(const api::PrimitivePtr &primitive) { auto stride_ptr = primitive->GetAttr(ops::kStride); if (stride_ptr != nullptr) { - auto stride_data = GetValue>(stride_ptr); + auto stride_data = api::GetValue>(stride_ptr); if (stride_data.size() > 1) { if (stride_data[0] < 1 || stride_data[0] > kMaxStrideSize || stride_data[1] < 1 || stride_data[1] > kMaxStrideSize) { @@ -99,10 +101,10 @@ bool CheckStride(const PrimitivePtr &primitive) { } return true; } -bool CheckDilation(const PrimitivePtr &primitive) { +bool CheckDilation(const api::PrimitivePtr &primitive) { auto dilation_ptr = primitive->GetAttr(ops::kDilation); if (dilation_ptr != nullptr) { - auto dilation_data = GetValue>(dilation_ptr); + auto dilation_data = api::GetValue>(dilation_ptr); if (dilation_data.size() > 1) { if (dilation_data[0] < 1 || dilation_data[0] > kMaxDilationSize || dilation_data[1] < 1 || dilation_data[1] > kMaxDilationSize) { @@ -114,15 +116,15 @@ bool CheckDilation(const PrimitivePtr &primitive) { } return true; } -bool CheckAttr(CNodePtr op, const PrimitivePtr &primitive, int64_t input_w) { +bool CheckAttr(const api::CNodePtr &op, const api::PrimitivePtr &primitive, int64_t input_w) { auto dilation_ptr = primitive->GetAttr(ops::kDilation); auto kernel_ptr = primitive->GetAttr(ops::kKernelSize); auto stride_ptr = primitive->GetAttr(ops::kStride); auto output_paddings_ptr = primitive->GetAttr(ops::kOutputPaddings); if (dilation_ptr != nullptr && kernel_ptr != nullptr) { - auto kernel_data = GetValue>(kernel_ptr); - auto dilation_data = GetValue>(dilation_ptr); + auto kernel_data = api::GetValue>(kernel_ptr); + auto dilation_data = api::GetValue>(dilation_ptr); if ((kernel_data[0] - 1) * dilation_data[0] + 1 > kMaxDilationAndKernelProd || (kernel_data[1] - 1) * dilation_data[1] + 1 > kMaxDilationAndKernelProd) { MS_LOG(WARNING) << "kernel should satisfy ((kernel - 1) * dilation + 1) less than 15"; @@ -130,15 +132,15 @@ bool CheckAttr(CNodePtr op, const PrimitivePtr &primitive, int64_t input_w) { } } if (stride_ptr != nullptr && kernel_ptr != nullptr) { - auto kernel_data = GetValue>(kernel_ptr); - auto stride_data = GetValue>(stride_ptr); - if (CheckPrimitiveType(op, prim::kPrimConv2DFusion)) { + auto kernel_data = api::GetValue>(kernel_ptr); + auto stride_data = api::GetValue>(stride_ptr); + if (CheckPrimitiveType(op, api::MakeShared())) { if (kernel_data[0] > kConvUpperBound / (input_w / (kCoefficient * stride_data[1]) * stride_data[1])) { MS_LOG(WARNING) << "kernel and stride should satisfy kernel_h <= 2048 / (w / (16 * stride) * stride) " << op->fullname_with_scope(); return false; } - } else if (CheckPrimitiveType(op, prim::kPrimConv2dTransposeFusion)) { + } else if (CheckPrimitiveType(op, api::MakeShared())) { MS_CHECK_TRUE_MSG(input_w != 0, false, "input_w should be 0."); if (kernel_data[0] > (kMaxNumOutput / input_w - 1.0) * stride_data[0] + 1.0) { MS_LOG(WARNING) << "kernel and stride should satisfy kernel_h <= (32768 / w - 1) * stride + 1.0 " @@ -148,8 +150,8 @@ bool CheckAttr(CNodePtr op, const PrimitivePtr &primitive, int64_t input_w) { } } - if (output_paddings_ptr != nullptr && CheckPrimitiveType(op, prim::kPrimConv2dTransposeFusion)) { - auto output_paddings = GetValue>(output_paddings_ptr); + if (output_paddings_ptr != nullptr && CheckPrimitiveType(op, api::MakeShared())) { + auto output_paddings = api::GetValue>(output_paddings_ptr); if (std::find_if_not(output_paddings.begin(), output_paddings.end(), [](int64_t output_padding) { return output_padding == 0; }) != output_paddings.end()) { MS_LOG(WARNING) << "output_padding attr only support 0 by dpico. " << op->fullname_with_scope(); @@ -160,8 +162,8 @@ bool CheckAttr(CNodePtr op, const PrimitivePtr &primitive, int64_t input_w) { } } // namespace -bool Conv2DFusionChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool Conv2DFusionChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.h index aa39a50816..f8a1422477 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/conv2d_checker.h @@ -26,7 +26,7 @@ class Conv2DFusionChecker : public OpChecker { public: Conv2DFusionChecker() : OpChecker("Conv2DFusion") {} ~Conv2DFusionChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.cc index 1db2cf79ad..adfc9b0f96 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.cc @@ -20,7 +20,7 @@ namespace mindspore { namespace dpico { -bool CustomOpChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { return true; } +bool CustomOpChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { return true; } OpCheckerRegistrar g_DecBboxChecker("DecBBox", new CustomOpChecker()); OpCheckerRegistrar g_DetectionOutputChecker("DetectionOutput", new CustomOpChecker()); OpCheckerRegistrar g_ExtractChecker("Extract", new CustomOpChecker()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.h index 7849b77b88..e65e5dc98f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/custom_op_checker.h @@ -26,7 +26,7 @@ class CustomOpChecker : public OpChecker { public: CustomOpChecker() : OpChecker("CustomOp") {} ~CustomOpChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.cc index c3eb2ffd69..5f8126f1bb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.cc @@ -23,16 +23,16 @@ namespace mindspore { namespace dpico { namespace { constexpr int kModeSize = 3; -} -bool EltwiseChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +} // namespace +bool EltwiseChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto mode_ptr = primitive->GetAttr(ops::kMode); if (mode_ptr != nullptr) { - auto mode_data = GetValue(mode_ptr); + auto mode_data = api::GetValue(mode_ptr); if (mode_data >= kModeSize) { // only prod(0), sum(1), max(2) is supported MS_LOG(WARNING) << "mode only supports 0/1/2 " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.h index 85b277c512..0690811313 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/eltwise_checker.h @@ -26,7 +26,7 @@ class EltwiseChecker : public OpChecker { public: EltwiseChecker() : OpChecker("Eltwise") {} ~EltwiseChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.cc index ca0b39eb65..b925894a0a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.cc @@ -17,26 +17,26 @@ #include "checker/exp_checker.h" #include #include -#include +#include #include "common/op_enum.h" namespace mindspore { namespace dpico { -bool ExpFusionChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ExpFusionChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, kInputIndex1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto base_ptr = primitive->GetAttr(ops::kBase); if (base_ptr != nullptr) { - auto base_data = GetValue(base_ptr); - if (base_data < 0 && fabs(base_data + 1.0) > std::numeric_limits::epsilon()) { + auto base_data = api::GetValue(base_ptr); + if (base_data < 0 && std::fabs(base_data + 1.0) > std::numeric_limits::epsilon()) { MS_LOG(WARNING) << "base val only supports -1 or positive num. " << op->fullname_with_scope(); return false; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.h index 7c7a08a1a8..cc586e032f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/exp_checker.h @@ -26,7 +26,7 @@ class ExpFusionChecker : public OpChecker { public: ExpFusionChecker() : OpChecker("ExpFusion") {} ~ExpFusionChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.cc index b93e0b765d..5471ccf698 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.cc @@ -24,22 +24,22 @@ namespace mindspore { namespace dpico { namespace { constexpr int kMaxFlattenInputW = 65536; -} -bool FlattenChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +} // namespace +bool FlattenChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr." << op->fullname_with_scope(); return false; } if (primitive->GetAttr(kStartAxis) != nullptr) { - auto axis = GetValue(primitive->GetAttr(kStartAxis)); + auto axis = api::GetValue(primitive->GetAttr(kStartAxis)); if (axis < kAxisLowerBound || axis > kAxisUpperBound) { MS_LOG(WARNING) << "start_axis val should in range [-4, 3] " << op->fullname_with_scope(); return false; } } if (primitive->GetAttr(kEndAxis) != nullptr) { - auto end_axis = GetValue(primitive->GetAttr(kEndAxis)); + auto end_axis = api::GetValue(primitive->GetAttr(kEndAxis)); if (end_axis < kAxisLowerBound || end_axis > kAxisUpperBound) { MS_LOG(WARNING) << "end_axis val should in range [-4, 3] " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.h index dcd41fe881..a3ac0f43b3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/flatten_checker.h @@ -26,7 +26,7 @@ class FlattenChecker : public OpChecker { public: FlattenChecker() : OpChecker("Flatten") {} ~FlattenChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.cc index 0d7ec51eb2..96fe614202 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.cc @@ -23,7 +23,7 @@ namespace mindspore { namespace dpico { -bool FullConnectionChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool FullConnectionChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; @@ -43,13 +43,13 @@ bool FullConnectionChecker::Check(CNodePtr op, int32_t output_num, mindspore::Fo return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } if (primitive->GetAttr(kNumOutput) != nullptr) { - auto num_output = GetValue(primitive->GetAttr(kNumOutput)); + auto num_output = api::GetValue(primitive->GetAttr(kNumOutput)); if (num_output > kMaxNumOutput) { MS_LOG(WARNING) << "num_output val:" << num_output << " should be less than " << kMaxNumOutput << " " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.h index 430fd15031..c69d521d48 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/full_connection_checker.h @@ -26,7 +26,7 @@ class FullConnectionChecker : public OpChecker { public: FullConnectionChecker() : OpChecker("FullConnection") {} ~FullConnectionChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.cc index 0f1ae25ffa..09eb4ba076 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.cc @@ -16,28 +16,28 @@ #include "checker/log_checker.h" #include -#include #include +#include #include "common/op_enum.h" namespace mindspore { namespace dpico { -bool LogChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool LogChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, kInputIndex1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto base_ptr = primitive->GetAttr(ops::kBase); if (base_ptr != nullptr) { - auto base_data = GetValue(base_ptr); // support -1.0 && any positive num but 1.0 - if (fabs(base_data + 1.0) <= std::numeric_limits::epsilon() || - (base_data > 0 && fabs(base_data - 1.0) > std::numeric_limits::epsilon())) { + auto base_data = api::GetValue(base_ptr); // support -1.0 && any positive num but 1.0 + if (std::fabs(base_data + 1.0) <= std::numeric_limits::epsilon() || + (base_data > 0 && std::fabs(base_data - 1.0) > std::numeric_limits::epsilon())) { return true; } else { MS_LOG(WARNING) << "base val only supports -1.0 or any positive num but 1.0 " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.h index 8e99c8304a..8ed99e64fb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/log_checker.h @@ -26,7 +26,7 @@ class LogChecker : public OpChecker { public: LogChecker() : OpChecker("Log") {} ~LogChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.cc index 83fe3e9572..2e139c2e36 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.cc @@ -26,20 +26,20 @@ constexpr int local_size_3 = 3; constexpr int local_size_5 = 5; constexpr int local_size_7 = 7; } // namespace -bool LRNChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool LRNChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto r_ptr = primitive->GetAttr(ops::kDepthRadius); if (r_ptr != nullptr) { - auto r_data = GetValue(r_ptr) * 2 + 1; + auto r_data = api::GetValue(r_ptr) * 2 + 1; if (r_data != local_size_3 && r_data != local_size_5 && r_data != local_size_7) { MS_LOG(WARNING) << "local size only supports 3/5/7 " << op->fullname_with_scope(); return false; @@ -47,7 +47,7 @@ bool LRNChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format } auto norm_region_ptr = primitive->GetAttr(ops::kNormRegion); if (norm_region_ptr != nullptr) { - auto norm_region_data = GetValue(norm_region_ptr); + auto norm_region_data = api::GetValue(norm_region_ptr); if (norm_region_data != "ACROSS_CHANNEL") { MS_LOG(WARNING) << "norm region only supports ACROSS_CHANNEL " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.h index 6d0d73c6fa..11e64047f9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/lrn_checker.h @@ -26,7 +26,7 @@ class LRNChecker : public OpChecker { public: LRNChecker() : OpChecker("LRN") {} ~LRNChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.cc index 0d7ce1d890..df137135ea 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.cc @@ -22,16 +22,17 @@ namespace mindspore { namespace dpico { namespace { constexpr uint32_t kLstmMaxNumOutput = 5456; -} -bool LstmChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +} // namespace + +bool LstmChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto num_outputs_ptr = primitive->GetAttr(dpico::kNumOutput); if (num_outputs_ptr != nullptr) { - auto num_outputs_data = GetValue(num_outputs_ptr); + auto num_outputs_data = static_cast(api::GetValue(num_outputs_ptr)); if (num_outputs_data > kLstmMaxNumOutput) { MS_LOG(WARNING) << "num_output should less than " << kLstmMaxNumOutput; return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.h index e85c401807..c3f5b34ce0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/lstm_checker.h @@ -26,7 +26,7 @@ class LstmChecker : public OpChecker { public: LstmChecker() : OpChecker("Lstm") {} ~LstmChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.cc index 2529f80214..f8131f9d7e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.cc @@ -24,7 +24,7 @@ namespace mindspore { namespace dpico { namespace { -bool CheckInputShapeForMatrix(const CNodePtr &cnode, const PrimitivePtr &primitive) { +bool CheckInputShapeForMatrix(const api::CNodePtr &cnode, const api::PrimitivePtr &primitive) { ShapeVector input_1_shape; ShapeVector input_2_shape; if (GetInputShapeFromCNode(cnode, kInputIndex1, &input_1_shape) != RET_OK || @@ -43,13 +43,13 @@ bool CheckInputShapeForMatrix(const CNodePtr &cnode, const PrimitivePtr &primiti } } } - primitive->AddAttr(kDim1, MakeValue(input_1_shape.at(input_1_shape.size() - kInputIndex2))); - primitive->AddAttr(kDim2, MakeValue(input_1_shape.at(input_1_shape.size() - 1))); - primitive->AddAttr(kDim3, MakeValue(input_2_shape.at(input_2_shape.size() - kInputIndex2))); + primitive->AddAttr(kDim1, api::MakeValue(input_1_shape.at(input_1_shape.size() - kInputIndex2))); + primitive->AddAttr(kDim2, api::MakeValue(input_1_shape.at(input_1_shape.size() - 1))); + primitive->AddAttr(kDim3, api::MakeValue(input_2_shape.at(input_2_shape.size() - kInputIndex2))); } return true; } -bool CheckInputShapeForFc(const CNodePtr &cnode) { +bool CheckInputShapeForFc(const api::CNodePtr &cnode) { ShapeVector input_shape; if (GetInputShapeFromCNode(cnode, kInputIndex1, &input_shape) == RET_OK) { return input_shape.size() == dpico::kDims2 && input_shape.at(0) == 1; @@ -57,18 +57,18 @@ bool CheckInputShapeForFc(const CNodePtr &cnode) { return false; } } // namespace -bool MatMulChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool MatMulChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr." << op->fullname_with_scope(); return false; } if (primitive->GetAttr(ops::kTransposeA) != nullptr) { - auto transpose_a = GetValue(primitive->GetAttr(ops::kTransposeA)); + auto transpose_a = api::GetValue(primitive->GetAttr(ops::kTransposeA)); if (transpose_a) { return false; } @@ -79,9 +79,9 @@ bool MatMulChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for } bool transpose_b = false; if (primitive->GetAttr(ops::kTransposeB) != nullptr) { - transpose_b = GetValue(primitive->GetAttr(ops::kTransposeB)); + transpose_b = api::GetValue(primitive->GetAttr(ops::kTransposeB)); } else { - primitive->AddAttr(ops::kTransposeB, MakeValue(false)); + primitive->AddAttr(ops::kTransposeB, api::MakeValue(false)); } if (!HasOfflineData(op->input(kInputIndex1))) { if (!HasOfflineData(op->input(kInputIndex2))) { @@ -89,14 +89,14 @@ bool MatMulChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for if (!CheckInputShapeForMatrix(op, primitive)) { return false; } - primitive->AddAttr(kOperatorType, MakeValue("Matrix")); + primitive->AddAttr(kOperatorType, api::MakeValue("Matrix")); return true; } else { return false; } } else { if (CheckInputShapeForFc(op)) { - primitive->AddAttr(kOperatorType, MakeValue("FullConnection")); + primitive->AddAttr(kOperatorType, api::MakeValue("FullConnection")); return true; } else { MS_LOG(WARNING) << "only supports input N = 1 by dpico. " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.h index 9005cac016..1ee20d3db6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/mat_mul_checker.h @@ -26,7 +26,7 @@ class MatMulChecker : public OpChecker { public: MatMulChecker() : OpChecker("MatMul") {} ~MatMulChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.cc index 797acb3112..3743810186 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.cc @@ -16,7 +16,6 @@ #include "checker/mvn_checker.h" #include -#include #include #include "common/anf_util.h" #include "common/op_enum.h" @@ -24,21 +23,21 @@ namespace mindspore { namespace dpico { -bool MvnChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool MvnChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, kInputIndex1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } if (primitive->GetAttr(ops::kAxes) != nullptr) { - auto axes = GetValue>(primitive->GetAttr(ops::kAxes)); - if (axes != std::vector{0, kAxis2, kAxis3} && axes != std::vector{0, 1, kAxis2, kAxis3}) { + auto axes = api::GetValue>(primitive->GetAttr(ops::kAxes)); + if (axes != std::vector{0, kAxis2, kAxis3} && axes != std::vector{0, 1, kAxis2, kAxis3}) { MS_LOG(WARNING) << "axes only supports [0,2,3] or [0,1,2,3] " << op->fullname_with_scope(); return false; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.h index 5d61f4d82b..8bb7e8c7ef 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/mvn_checker.h @@ -26,7 +26,7 @@ class MvnChecker : public OpChecker { public: MvnChecker() : OpChecker("Mvn") {} ~MvnChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.cc index 62578adf91..227bcfbd60 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.cc @@ -99,12 +99,12 @@ STATUS GetVectorChannel(const std::vector &shape, int64_t *channel) { return RET_OK; } -bool HasOfflineData(const AnfNodePtr &node) { - auto param = node->cast(); +bool HasOfflineData(const api::AnfNodePtr &node) { + auto param = node->cast(); return param != nullptr && param->has_default(); } -bool CheckInputW(const CNodePtr &op, size_t index, mindspore::Format format, int limit_w) { +bool CheckInputW(const api::CNodePtr &op, size_t index, mindspore::Format format, int limit_w) { if (index >= op->inputs().size()) { MS_LOG(ERROR) << "index:" << index << " is greater than " << op->fullname_with_scope() << " inputs size:" << op->inputs().size(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.h index fb30d83209..6d8fe17bb5 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/op_checker.h @@ -23,9 +23,12 @@ #include #include #include +#include "mindapi/base/format.h" #include "include/errorcode.h" -#include "ir/anf.h" -#include "ops/op_utils.h" +#include "mindapi/ir/anf.h" +#include "mindapi/base/logging.h" +#include "ops/op_name.h" + using mindspore::lite::RET_ERROR; using mindspore::lite::RET_OK; using mindspore::lite::STATUS; @@ -35,7 +38,7 @@ class OpChecker { public: explicit OpChecker(std::string node_name) : name(std::move(node_name)) {} virtual ~OpChecker() = default; - virtual bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) = 0; + virtual bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) = 0; protected: const std::string name; @@ -44,8 +47,8 @@ class OpChecker { STATUS GetWidth(const std::vector &shape, mindspore::Format format, int64_t *width); STATUS GetTensorChannel(const std::vector &shape, mindspore::Format format, int64_t *channel); STATUS GetVectorChannel(const std::vector &shape, int64_t *channel); -bool HasOfflineData(const AnfNodePtr &node); -bool CheckInputW(const CNodePtr &op, size_t index, mindspore::Format format, int limit_w); +bool HasOfflineData(const api::AnfNodePtr &node); +bool CheckInputW(const api::CNodePtr &op, size_t index, mindspore::Format format, int limit_w); class OpCheckerRegistry { public: diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.cc index 9d261d3d45..1b9a5f185b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.cc @@ -33,10 +33,10 @@ constexpr int kMaxDilationSize = 5; constexpr int kMaxDilationAndKernelProd = 15; constexpr float kCoefficient = 16.0; -bool CheckPad(const PrimitivePtr &primitive) { +bool CheckPad(const api::PrimitivePtr &primitive) { auto pad_ptr = primitive->GetAttr(ops::kPad); if (pad_ptr != nullptr) { - auto pad_data = GetValue>(pad_ptr); + auto pad_data = api::GetValue>(pad_ptr); if (pad_data.size() > kDims3) { if (pad_data[0] < 0 || pad_data[0] > kMaxPadSize || pad_data[1] < 0 || pad_data[1] > kMaxPadSize || pad_data[kAxis2] < 0 || pad_data[kAxis2] > kMaxPadSize || pad_data[kAxis3] < 0 || @@ -48,13 +48,13 @@ bool CheckPad(const PrimitivePtr &primitive) { } return true; } -bool CheckKernelSize(const PrimitivePtr &primitive) { +bool CheckKernelSize(const api::PrimitivePtr &primitive) { auto is_global_ptr = primitive->GetAttr(ops::kGlobal); auto kernel_ptr = primitive->GetAttr(ops::kKernelSize); if (kernel_ptr != nullptr) { - auto kernel_data = GetValue>(kernel_ptr); + auto kernel_data = api::GetValue>(kernel_ptr); if (is_global_ptr != nullptr) { - auto is_global = GetValue(is_global_ptr); + auto is_global = api::GetValue(is_global_ptr); if (!is_global && (kernel_data[0] < 1 || kernel_data[0] > kMaxKernelSize || kernel_data[1] < 1 || kernel_data[1] > kMaxKernelSize)) { MS_LOG(WARNING) << "kernel should in range [1,255]"; @@ -64,10 +64,10 @@ bool CheckKernelSize(const PrimitivePtr &primitive) { } return true; } -bool CheckStride(const PrimitivePtr &primitive) { +bool CheckStride(const api::PrimitivePtr &primitive) { auto stride_ptr = primitive->GetAttr(ops::kStrides); if (stride_ptr != nullptr) { - auto stride_data = GetValue>(stride_ptr); + auto stride_data = api::GetValue>(stride_ptr); if (stride_data[0] < 1 || stride_data[0] > kMaxStrideSize || stride_data[1] < 1 || stride_data[1] > kMaxStrideSize) { MS_LOG(WARNING) << "stride should in range [1,255]"; @@ -76,10 +76,10 @@ bool CheckStride(const PrimitivePtr &primitive) { } return true; } -bool CheckRoundMode(const PrimitivePtr &primitive) { +bool CheckRoundMode(const api::PrimitivePtr &primitive) { auto round_mode_ptr = primitive->GetAttr(ops::kRoundMode); if (round_mode_ptr != nullptr) { - auto round_mode_data = GetValue(round_mode_ptr); + auto round_mode_data = api::GetValue(round_mode_ptr); if (round_mode_data != 0 && round_mode_data != 1) { MS_LOG(WARNING) << "round mode only supports CEIL or FLOOR"; return false; @@ -87,12 +87,12 @@ bool CheckRoundMode(const PrimitivePtr &primitive) { } return true; } -bool CheckAttr(const PrimitivePtr &primitive, int64_t input_w) { +bool CheckAttr(const api::PrimitivePtr &primitive, int64_t input_w) { auto kernel_ptr = primitive->GetAttr(ops::kKernelSize); auto stride_ptr = primitive->GetAttr(ops::kStride); if (stride_ptr != nullptr && kernel_ptr != nullptr) { - auto kernel_data = GetValue>(kernel_ptr); - auto stride_data = GetValue>(stride_ptr); + auto kernel_data = api::GetValue>(kernel_ptr); + auto stride_data = api::GetValue>(stride_ptr); if (kernel_data[0] > kMaxPoolUpperBound / (input_w / (kCoefficient * stride_data[1]) * stride_data[1])) { MS_LOG(WARNING) << "kernel and stride should satisfy kernel_h <= 2048 / (w / (16 * stride) * stride) "; return false; @@ -102,8 +102,8 @@ bool CheckAttr(const PrimitivePtr &primitive, int64_t input_w) { } } // namespace -bool PoolingChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool PoolingChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.h index 6d08a0a813..c07104a11e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/pooling_checker.h @@ -26,7 +26,7 @@ class PoolingChecker : public OpChecker { public: PoolingChecker() : OpChecker("Pooling") {} ~PoolingChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.cc index ceb98a986c..305722e759 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.cc @@ -17,19 +17,19 @@ #include "checker/pow_checker.h" #include #include -#include +#include #include "common/fetch_content.h" #include "common/op_enum.h" namespace mindspore { namespace dpico { -bool PowFusionChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool PowFusionChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, kInputIndex1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } float power; - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; @@ -46,13 +46,14 @@ bool PowFusionChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format } power = *(reinterpret_cast(data_info.data_.data())); } else if (primitive->GetAttr(ops::kPower) != nullptr) { - power = GetValue(primitive->GetAttr(ops::kPower)); + power = api::GetValue(primitive->GetAttr(ops::kPower)); } else { MS_LOG(ERROR) << "null param"; return false; } - if (!(fmod(fabs(power), 1.0) > std::numeric_limits::epsilon() && // support power: -0.5, 0.5, integers - fabs(fabs(power) - 0.5) > std::numeric_limits::epsilon())) { + if (!(std::fmod(std::fabs(power), 1.0) > + std::numeric_limits::epsilon() && // support power: -0.5, 0.5, integers + std::fabs(std::fabs(power) - 0.5) > std::numeric_limits::epsilon())) { return true; } else { MS_LOG(WARNING) << "power val only supports -0.5, 0.5, integers " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.h index c2750e2c1a..4dcd4588e8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/pow_checker.h @@ -26,7 +26,7 @@ class PowFusionChecker : public OpChecker { public: PowFusionChecker() : OpChecker("PowFusion") {} ~PowFusionChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.cc index 0eb580f3ef..3b12059928 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.cc @@ -18,6 +18,7 @@ #include #include #include +#include "mindapi/base/types.h" #include "include/registry/converter_context.h" #include "common/fetch_content.h" #include "common/anf_util.h" @@ -26,7 +27,7 @@ namespace mindspore { namespace dpico { namespace { -STATUS GetAxesSet(const CNodePtr &op, const ShapeVector &input_shape, const PrimitivePtr &primitive, +STATUS GetAxesSet(const api::CNodePtr &op, const ShapeVector &input_shape, const api::PrimitivePtr &primitive, std::set *axes_set) { if (axes_set == nullptr) { MS_LOG(ERROR) << "axes_set is nullptr. " << op->fullname_with_scope(); @@ -51,18 +52,19 @@ STATUS GetAxesSet(const CNodePtr &op, const ShapeVector &input_shape, const Prim (void)std::transform(data, data + data_size, std::inserter(*axes_set, axes_set->begin()), [input_shape](int32_t value) { return (value + input_shape.size()) % input_shape.size(); }); } else if (primitive->GetAttr(ops::kAxes) != nullptr) { - auto axes = GetValue>(primitive->GetAttr(ops::kAxes)); - (void)std::transform(axes.begin(), axes.end(), std::inserter(*axes_set, axes_set->begin()), - [input_shape](int32_t value) { return (value + input_shape.size()) % input_shape.size(); }); + auto axes = api::GetValue>(primitive->GetAttr(ops::kAxes)); + (void)std::transform( + axes.begin(), axes.end(), std::inserter(*axes_set, axes_set->begin()), + [input_shape](int64_t value) { return (static_cast(value) + input_shape.size()) % input_shape.size(); }); } return RET_OK; } -bool CheckAttr(const CNodePtr &op, mindspore::Format format, const PrimitivePtr &primitive, +bool CheckAttr(const api::CNodePtr &op, mindspore::Format format, const api::PrimitivePtr &primitive, const ShapeVector &input_shape) { bool keep_dims = true; if (primitive->GetAttr(ops::kKeepDims) != nullptr) { - keep_dims = GetValue(primitive->GetAttr(ops::kKeepDims)); + keep_dims = api::GetValue(primitive->GetAttr(ops::kKeepDims)); } std::set axes_set; if (GetAxesSet(op, input_shape, primitive, &axes_set) != RET_OK) { @@ -71,7 +73,7 @@ bool CheckAttr(const CNodePtr &op, mindspore::Format format, const PrimitivePtr } // special process when cnode is from caffe if (primitive->GetAttr(ops::kFmkType) != nullptr) { - auto fmk_type = static_cast(GetValue(primitive->GetAttr(ops::kFmkType))); + auto fmk_type = static_cast(api::GetValue(primitive->GetAttr(ops::kFmkType))); if (fmk_type == converter::kFmkTypeCaffe) { return true; } @@ -86,15 +88,15 @@ bool CheckAttr(const CNodePtr &op, mindspore::Format format, const PrimitivePtr return true; } } // namespace -bool ReduceChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool ReduceChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto mode_ptr = primitive->GetAttr(ops::kMode); if (mode_ptr != nullptr) { - auto reduce_mode = static_cast(GetValue(mode_ptr)); + auto reduce_mode = static_cast(api::GetValue(mode_ptr)); if (reduce_mode < ReduceMode::Reduce_Mean || reduce_mode >= ReduceMode::Reduce_All) { MS_LOG(WARNING) << "unsupported reduce mode " << reduce_mode << " " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.h index e49a7c07f1..0d0f55ca86 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/reduce_checker.h @@ -26,7 +26,7 @@ class ReduceChecker : public OpChecker { public: ReduceChecker() : OpChecker("Reduce") {} ~ReduceChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.cc index 8169c7f29a..1ccf523621 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.cc @@ -27,7 +27,7 @@ namespace dpico { namespace { constexpr int kMaxReshapeInputW = 65536; } // namespace -bool ReshapeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ReshapeChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { std::vector input_shape; if (GetInputShapeFromCNode(op, kInputIndex1, &input_shape) == RET_OK && !input_shape.empty()) { int64_t input_w; @@ -40,7 +40,7 @@ bool ReshapeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format fo return false; } } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; @@ -53,7 +53,7 @@ bool ReshapeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format fo DataInfo data_info; std::vector shape_data; - auto shape_ptr = primitive->GetAttr(kShape); + auto shape_ptr = primitive->GetAttr(ops::kShape); if (op->inputs().size() > kInputIndex2 && FetchDataFromParameterNode(op, kInputIndex2, &data_info) == lite::RET_OK) { if (data_info.data_type_ != kNumberTypeInt32) { MS_LOG(ERROR) << "data_type not correct"; @@ -72,19 +72,19 @@ bool ReshapeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format fo (void)std::transform(data, data + data_size, std::back_inserter(shape_data), [](const int32_t &value) { return static_cast(value); }); } else if (shape_ptr != nullptr) { - shape_data = GetValue>(shape_ptr); + shape_data = api::GetValue>(shape_ptr); } else { MS_LOG(ERROR) << "can't get shape value. " << op->fullname_with_scope(); return false; } - primitive->AddAttr(kShape, MakeValue(shape_data)); + primitive->AddAttr(ops::kShape, api::MakeValue(shape_data)); - auto param_ptr = op->input(kInputIndex2)->cast(); + auto param_ptr = op->input(kInputIndex2)->cast(); if (param_ptr == nullptr) { MS_LOG(ERROR) << "param_ptr is nullptr. " << op->fullname_with_scope(); return false; } - auto param_value = std::dynamic_pointer_cast(param_ptr->default_param()); + auto param_value = param_ptr->default_param()->cast(); if (param_value == nullptr) { MS_LOG(ERROR) << "param_value is nullptr." << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.h index 333f5b5df1..87ab2d6c6e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/reshape_checker.h @@ -26,7 +26,7 @@ class ReshapeChecker : public OpChecker { public: ReshapeChecker() : OpChecker("Reshape") {} ~ReshapeChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.cc index 31824a6240..12e3fd9276 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.cc @@ -21,24 +21,25 @@ #include "common/check_base.h" #include "common/op_enum.h" #include "common/anf_util.h" +#include "mindapi/base/types.h" #include "include/registry/converter_context.h" namespace mindspore { namespace dpico { namespace { constexpr int kMaxOutputWOf4Dims = 2048; -bool IsFromCaffe(const PrimitivePtr &primitive) { +bool IsFromCaffe(const api::PrimitivePtr &primitive) { if (primitive->GetAttr(ops::kFmkType) != nullptr) { - auto fmk_type = static_cast(GetValue(primitive->GetAttr(ops::kFmkType))); + auto fmk_type = static_cast(api::GetValue(primitive->GetAttr(ops::kFmkType))); if (fmk_type == converter::kFmkTypeCaffe) { return true; } } return false; } -bool CheckInterpOp(const CNodePtr &op, const ShapeVector &input_shape, const ShapeVector &output_shape, size_t index_h, - size_t index_w) { - auto prim = GetValueNode(op->input(0)); +bool CheckInterpOp(const api::CNodePtr &op, const ShapeVector &input_shape, const ShapeVector &output_shape, + size_t index_h, size_t index_w) { + auto prim = api::GetValueNode(op->input(0)); MS_CHECK_TRUE_MSG(prim != nullptr, false, "prim is nullptr."); if (input_shape.at(index_w) > kMaxInputWOf4Dims || output_shape.at(index_w) >= kMaxOutputWOf4Dims) { MS_LOG(WARNING) << op->fullname_with_scope() << "'s input_w should be less than " << kMaxInputWOf4Dims @@ -48,10 +49,10 @@ bool CheckInterpOp(const CNodePtr &op, const ShapeVector &input_shape, const Sha int pad_beg = 0; int pad_end = 0; if (prim->GetAttr(dpico::kPadBeg) != nullptr) { - pad_beg = GetValue(prim->GetAttr(dpico::kPadBeg)); + pad_beg = api::GetValue(prim->GetAttr(dpico::kPadBeg)); } if (prim->GetAttr(dpico::kPadEnd) != nullptr) { - pad_end = GetValue(prim->GetAttr(dpico::kPadEnd)); + pad_end = api::GetValue(prim->GetAttr(dpico::kPadEnd)); } if (pad_beg > 0 || pad_end > 0) { MS_LOG(WARNING) << "pad_beg or pad_end only supports non-negative integer by dpico. " << op->fullname_with_scope(); @@ -61,19 +62,19 @@ bool CheckInterpOp(const CNodePtr &op, const ShapeVector &input_shape, const Sha return ((input_shape.at(index_h) + pad_beg + pad_end == 1) ^ (input_shape.at(index_w) + pad_beg + pad_end != 1)) && ((output_shape.at(index_h) == 1) ^ (output_shape.at(index_w) != 1)); } -bool IsDoubleResize(const CNodePtr &cnode, const ShapeVector &input_shape, const ShapeVector &output_shape, +bool IsDoubleResize(const api::CNodePtr &cnode, const ShapeVector &input_shape, const ShapeVector &output_shape, size_t index_h, size_t index_w) { const int64_t nums2 = 2; return input_shape.at(index_h) * nums2 == output_shape.at(index_h) && input_shape.at(index_w) * nums2 == output_shape.at(index_w); } -bool CheckResizeOp(const CNodePtr &op, const ShapeVector &input_shape, const ShapeVector &output_shape, size_t index_h, - size_t index_w) { - auto prim = GetValueNode(op->input(0)); +bool CheckResizeOp(const api::CNodePtr &op, const ShapeVector &input_shape, const ShapeVector &output_shape, + size_t index_h, size_t index_w) { + auto prim = api::GetValueNode(op->input(0)); MS_CHECK_TRUE_MSG(prim != nullptr, false, "prim is nullptr."); if (prim->GetAttr(ops::kCoordinateTransformMode) != nullptr) { auto coordinate_transform_mode = - static_cast(GetValue(prim->GetAttr(ops::kCoordinateTransformMode))); + static_cast(api::GetValue(prim->GetAttr(ops::kCoordinateTransformMode))); if (coordinate_transform_mode != CoordinateTransformMode::ASYMMETRIC) { MS_LOG(WARNING) << "resize only supports CoordinateTransformMode::ASYMMETRIC by dpico. " << op->fullname_with_scope(); @@ -81,14 +82,14 @@ bool CheckResizeOp(const CNodePtr &op, const ShapeVector &input_shape, const Sha } } if (prim->GetAttr(ops::kMethod) != nullptr) { - auto interpolation_mode = static_cast(GetValue(prim->GetAttr(ops::kMethod))); + auto interpolation_mode = static_cast(api::GetValue(prim->GetAttr(ops::kMethod))); if (interpolation_mode != ResizeMethod::NEAREST) { MS_LOG(WARNING) << "resize only supports ResizeMethod::NEAREST by dpico. " << op->fullname_with_scope(); return false; } } if (prim->GetAttr(ops::kNearestMode) != nullptr) { - auto nearest_mode = static_cast(GetValue(prim->GetAttr(ops::kNearestMode))); + auto nearest_mode = static_cast(api::GetValue(prim->GetAttr(ops::kNearestMode))); if (nearest_mode != mindspore::NearestMode::FLOOR) { MS_LOG(WARNING) << "resize only supports NearestMode::FLOOR by dpico. " << op->fullname_with_scope(); return false; @@ -104,8 +105,8 @@ bool CheckResizeOp(const CNodePtr &op, const ShapeVector &input_shape, const Sha return true; } } // namespace -bool ResizeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool ResizeChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); MS_CHECK_TRUE_MSG(primitive != nullptr, false, "prim is nullptr."); ShapeVector input_shape; auto abstract = GetCNodeInputAbstract(op, 1); @@ -135,7 +136,7 @@ bool ResizeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format for Format input_format; if (primitive->GetAttr(ops::kFormat) != nullptr) { - input_format = static_cast(GetValue(primitive->GetAttr(ops::kFormat))); + input_format = static_cast(api::GetValue(primitive->GetAttr(ops::kFormat))); } else { MS_LOG(ERROR) << ops::kFormat << " attr is needed."; return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.h index 63d5d5edbe..345774899f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/resize_checker.h @@ -26,7 +26,7 @@ class ResizeChecker : public OpChecker { public: ResizeChecker() : OpChecker("Resize") {} ~ResizeChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.cc index 6a52c2ca2a..ff720f7523 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.cc @@ -22,19 +22,19 @@ namespace mindspore { namespace dpico { -bool ReverseChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ReverseChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, kInputIndex1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } if (primitive->GetAttr(ops::kAxis) != nullptr) { - auto axis = GetValue>(primitive->GetAttr(ops::kAxis)); + auto axis = api::GetValue>(primitive->GetAttr(ops::kAxis)); if (axis.size() != 1) { MS_LOG(WARNING) << "reverse's axis size only supports 1 by dpico. " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.h index c0c695358e..1cb001e3e1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/reverse_checker.h @@ -26,7 +26,7 @@ class ReverseChecker : public OpChecker { public: ReverseChecker() : OpChecker("Reverse") {} ~ReverseChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.cc index 80ff6288b3..9d6da8b4d9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.cc @@ -26,20 +26,20 @@ namespace { constexpr int kNegativeAxisCorrespondZero = -4; constexpr int kNegativeAxisCorrespondOne = -3; } // namespace -bool ScaleChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool ScaleChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, 1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto axis_ptr = primitive->GetAttr(ops::kAxis); if (axis_ptr != nullptr) { - auto axis_data = GetValue(axis_ptr); + auto axis_data = api::GetValue(axis_ptr); std::unordered_set range = {1, kNegativeAxisCorrespondOne, 0, kNegativeAxisCorrespondZero}; if (range.find(axis_data) == range.end()) { MS_LOG(WARNING) << "axis val only supports 1/-3/0/-4 by dpico. " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.h index 9e61871f3a..f80c2c467b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/scale_checker.h @@ -26,7 +26,7 @@ class ScaleChecker : public OpChecker { public: ScaleChecker() : OpChecker("Scale") {} ~ScaleChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.cc index a3541cf230..900f4d3cdc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.cc @@ -23,22 +23,22 @@ namespace mindspore { namespace dpico { namespace { constexpr int kMaxSplitSize = 31; -} +} // namespace -bool SliceChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool SliceChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (output_num > kMaxTopNum) { MS_LOG(WARNING) << "output num " << output_num << " is greater than " << kMaxNumOutput << " " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } auto axis_ptr = primitive->GetAttr(ops::kAxes); if (axis_ptr != nullptr) { - auto axis_data = GetValue>(axis_ptr); + auto axis_data = api::GetValue>(axis_ptr); if (axis_data[0] < kAxisLowerBound || axis_data[0] > kAxisUpperBound) { MS_LOG(WARNING) << "axis val should in range [-4, 3]. " << op->fullname_with_scope(); return false; @@ -46,7 +46,7 @@ bool SliceChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format form } if (primitive->GetAttr(ops::kSizeSplits) != nullptr) { - auto splits = GetValue>(primitive->GetAttr(ops::kSizeSplits)); + auto splits = api::GetValue>(primitive->GetAttr(ops::kSizeSplits)); if (splits.size() > kMaxSplitSize) { MS_LOG(WARNING) << "split size should be less than " << kMaxSplitSize << " " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.h index 79d40ec075..4342b6a8f5 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/slice_checker.h @@ -26,7 +26,7 @@ class SliceChecker : public OpChecker { public: SliceChecker() : OpChecker("Slice") {} ~SliceChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.cc index 55a7286c7b..a0c6a549fd 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.cc @@ -48,19 +48,19 @@ bool CheckVectorAndTensorChannel(const std::vector &input_shape, mindsp return true; } } // namespace -bool SoftmaxChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool SoftmaxChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { std::vector input_shape; if (GetInputShapeFromCNode(op, kInputIndex1, &input_shape) == RET_OK && !input_shape.empty()) { - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); MS_CHECK_TRUE_MSG(primitive != nullptr, false, "primitive is nullptr."); auto axis_ptr = primitive->GetAttr(ops::kAxis); if (axis_ptr != nullptr) { - auto axis_data = GetValue>(axis_ptr); + auto axis_data = api::GetValue>(axis_ptr); auto axis = axis_data[0]; if (axis < 0) { axis = (axis + input_shape.size()) % input_shape.size(); std::vector axes = {axis}; - primitive->set_attr(ops::kAxis, MakeValue(axes)); + primitive->AddAttr(ops::kAxis, api::MakeValue(axes)); } if (axis != kAxis1 && axis != kAxis2 && axis != kAxis3) { MS_LOG(WARNING) << "axis val only supports 1/2/3 by dpico. " << op->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.h index ca7a9e7f3e..9cc79b0767 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/softmax_checker.h @@ -26,7 +26,7 @@ class SoftmaxChecker : public OpChecker { public: SoftmaxChecker() : OpChecker("Softmax") {} ~SoftmaxChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.cc index 272a93219b..55124a92d3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.cc @@ -20,16 +20,16 @@ namespace mindspore { namespace dpico { -bool SplitChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool SplitChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (output_num > kMaxTopNum) { MS_LOG(WARNING) << op->fullname_with_scope() << "'s output num " << output_num << " is greater than " << kMaxTopNum; return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { return false; } - return primitive->GetAttr(ops::kAxis) != nullptr && GetValue(primitive->GetAttr(ops::kAxis)) != 0; + return primitive->GetAttr(ops::kAxis) != nullptr && api::GetValue(primitive->GetAttr(ops::kAxis)) != 0; } OpCheckerRegistrar g_SplitChecker("Split", new SplitChecker()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.h index 14703c1aae..f6e559f5d9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/split_checker.h @@ -26,7 +26,7 @@ class SplitChecker : public OpChecker { public: SplitChecker() : OpChecker("Split") {} ~SplitChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.cc index 4ec63a63c7..171fbb1ae4 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.cc @@ -20,14 +20,14 @@ namespace mindspore { namespace dpico { -bool SppChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool SppChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } if (primitive->GetAttr(dpico::kPoolMethod) != nullptr) { - auto pool_method = GetValue(primitive->GetAttr(dpico::kPoolMethod)); + auto pool_method = api::GetValue(primitive->GetAttr(dpico::kPoolMethod)); if (pool_method != 0 && pool_method != 1) { MS_LOG(WARNING) << "only supports max && ave pool by dpico. " << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.h index 8b56a347a8..931dc12d6d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/spp_checker.h @@ -26,7 +26,7 @@ class SppChecker : public OpChecker { public: SppChecker() : OpChecker("Spp") {} ~SppChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.cc index 29110eb56e..3ecd80fb20 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace dpico { -bool SqueezeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool SqueezeChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { std::vector output_shapes; if (GetBoolAttr(op, dpico::kInferDone)) { if (GetOutputShapesFromCNode(op, &output_shapes) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.h index ed5e50b3de..e9cf3e42e2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/squeeze_checker.h @@ -26,7 +26,7 @@ class SqueezeChecker : public OpChecker { public: SqueezeChecker() : OpChecker("Squeeze") {} ~SqueezeChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.cc index ff235fca3a..42804bb362 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.cc @@ -21,14 +21,14 @@ namespace mindspore { namespace dpico { -bool StridedSliceChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { - auto primitive = GetValueNode(op->input(0)); +bool StridedSliceChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; } if (primitive->GetAttr(ops::kFmkType) != nullptr) { - auto fmk_type = static_cast(GetValue(primitive->GetAttr(ops::kFmkType))); + auto fmk_type = static_cast(api::GetValue(primitive->GetAttr(ops::kFmkType))); if (fmk_type == converter::kFmkTypeOnnx) { return true; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.h index 608baf90ac..350426ec1f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/strided_slice_checker.h @@ -26,7 +26,7 @@ class StridedSliceChecker : public OpChecker { public: StridedSliceChecker() : OpChecker("StridedSlice") {} ~StridedSliceChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.cc b/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.cc index 221f5c98f1..0fbc3e41e3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.cc @@ -17,6 +17,7 @@ #include "checker/transpose_checker.h" #include #include +#include #include "common/data_transpose_utils.h" #include "common/fetch_content.h" #include "common/op_attr.h" @@ -24,12 +25,12 @@ namespace mindspore { namespace dpico { -bool TransposeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format format) { +bool TransposeChecker::Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) { if (!CheckInputW(op, kInputIndex1, format, kMaxInputWOf4Dims)) { MS_LOG(WARNING) << "input_w is not supported. " << op->fullname_with_scope(); return false; } - auto primitive = GetValueNode(op->input(0)); + auto primitive = api::GetValueNode(op->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr"; return false; @@ -52,7 +53,9 @@ bool TransposeChecker::Check(CNodePtr op, int32_t output_num, mindspore::Format } perm_val = {data[0], data[kAxis1], data[kAxis2], data[kAxis3]}; } else if (primitive->GetAttr(kPerm) != nullptr) { - perm_val = GetValue>(primitive->GetAttr(kPerm)); + auto perm_vec = api::GetValue>(primitive->GetAttr(kPerm)); + (void)std::transform(perm_vec.begin(), perm_vec.end(), std::back_inserter(perm_val), + [](int64_t p) { return static_cast(p); }); } else { MS_LOG(ERROR) << "transpose param invalid" << op->fullname_with_scope(); return false; diff --git a/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.h b/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.h index 0033f4b4b4..77b80b4157 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.h +++ b/mindspore/lite/tools/converter/adapter/dpico/checker/transpose_checker.h @@ -26,7 +26,7 @@ class TransposeChecker : public OpChecker { public: TransposeChecker() : OpChecker("Transpose") {} ~TransposeChecker() override = default; - bool Check(CNodePtr op, int32_t output_num, mindspore::Format format) override; + bool Check(api::CNodePtr op, int32_t output_num, mindspore::Format format) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.cc b/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.cc index 2da670bf76..d8905de4a9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.cc @@ -19,117 +19,126 @@ #include #include #include +#include +#include +#include "third_party/securec/include/securec.h" #include "common/op_enum.h" #include "common/op_attr.h" #include "common/string_util.h" #include "ops/custom.h" +#include "ops/tuple_get_item.h" #include "ops/transpose.h" -#include "nnacl/op_base.h" +#include "common/check_base.h" namespace mindspore { namespace dpico { namespace { +const std::map kTypeMap = { + {kNumberTypeBool, 1}, {kNumberTypeInt, 4}, {kNumberTypeInt8, 1}, {kNumberTypeInt16, 2}, + {kNumberTypeInt32, 4}, {kNumberTypeInt64, 8}, {kNumberTypeUInt, 4}, {kNumberTypeUInt8, 1}, + {kNumberTypeUInt16, 2}, {kNumberTypeUInt32, 4}, {kNumberTypeUInt64, 8}, {kNumberTypeFloat, 4}, + {kNumberTypeFloat16, 2}, {kNumberTypeFloat32, 4}, {kNumberTypeFloat64, 8}, {kNumberTypeComplex64, 8}, + {kNumberTypeComplex128, 16}}; constexpr size_t kTupleGetItemInputSize = 3; constexpr size_t kInputNodeOutputIndexInTupleGetItem = 2; using PrimitiveCPtr = std::shared_ptr; +size_t TypeIdSize(const TypeId data_type) { + const size_t unsupported_type_error = 0; + auto iter = kTypeMap.find(data_type); + if (iter != kTypeMap.end()) { + return iter->second; + } + return unsupported_type_error; +} } // namespace -bool CheckPrimitiveType(const AnfNodePtr &node, const PrimitivePtr &primitive_type) { +bool CheckPrimitiveType(const api::AnfNodePtr &node, const api::PrimitivePtr &primitive_type) { if (node == nullptr) { return false; } - if (node->isa()) { - auto cnode = node->cast(); + if (node->isa()) { + auto cnode = node->cast(); return IsPrimitive(cnode->input(0), primitive_type); - } else if (node->isa()) { + } else if (node->isa()) { return IsPrimitive(node, primitive_type); } return false; } -STATUS GetPrimitiveType(const AnfNodePtr &node, std::string *name) { +STATUS GetPrimitiveType(const api::AnfNodePtr &node, std::string *name) { if (name == nullptr) { MS_LOG(ERROR) << "name is nulltr."; return RET_ERROR; } - if (node->isa()) { - auto cnode = node->cast(); - auto primitive = GetValueNode(cnode->input(0)); + if (node->isa()) { + auto cnode = node->cast(); + auto primitive = api::GetValueNode(cnode->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr. " << cnode->fullname_with_scope(); return RET_ERROR; } - if (CheckPrimitiveType(node, prim::kPrimCustom)) { - auto custom_prim = utils::cast>(primitive); - MS_ASSERT(custom_prim != nullptr); + if (CheckPrimitiveType(node, api::MakeShared())) { + auto custom_prim = api::utils::cast>(primitive); + MS_CHECK_TRUE_MSG(custom_prim != nullptr, RET_ERROR, "custom op is nullptr."); *name = custom_prim->get_type(); return RET_OK; } else { *name = primitive->name(); return RET_OK; } - } else if (node->isa()) { - auto fn_value = GetValueNode(node); + } else if (node->isa()) { + auto fn_value = api::GetValueNode(node); *name = fn_value->name(); return RET_OK; } MS_LOG(ERROR) << "There is no name for this node"; return RET_ERROR; } -STATUS GetShapeVectorFromParameter(const AnfNodePtr &anode, ShapeVector *shape_vector) { +STATUS GetShapeVectorFromParameter(const api::AnfNodePtr &anode, ShapeVector *shape_vector) { if (shape_vector == nullptr) { MS_LOG(ERROR) << "shape vector is nullptr."; return RET_ERROR; } - if (!utils::isa(anode)) { + if (!api::utils::isa(anode)) { MS_LOG(ERROR) << "anode should be parameter node. "; return RET_ERROR; } - auto param_node = anode->cast(); + auto param_node = anode->cast(); auto abstract_base = param_node->abstract(); if (abstract_base == nullptr) { MS_LOG(ERROR) << "Abstract of parameter is nullptr, " << param_node->name(); return lite::RET_PARAM_INVALID; } - if (!utils::isa(abstract_base)) { + if (!api::utils::isa(abstract_base)) { MS_LOG(ERROR) << "Abstract of parameter should be abstract tensor, " << param_node->name(); return lite::RET_INPUT_TENSOR_ERROR; } - auto abstract_tensor = utils::cast(abstract_base); + auto abstract_tensor = api::utils::cast(abstract_base); MS_CHECK_TRUE_MSG(abstract_tensor != nullptr, RET_ERROR, "Cast to abstract tensor failed!"); - if (!utils::isa(abstract_tensor->BuildShape())) { + if (!api::utils::isa(abstract_tensor->shape())) { MS_LOG(ERROR) << "Shape of Abstract of parameter should be ShapePtr, " << param_node->name(); return lite::RET_PARAM_INVALID; } - *shape_vector = utils::cast(abstract_tensor->BuildShape())->shape(); + *shape_vector = api::utils::cast(abstract_tensor->shape())->shape(); return RET_OK; } -std::vector CastToInt(const ValuePtr &value) { +std::vector CastToInt(const api::ValuePtr &value) { if (value == nullptr) { MS_LOG(WARNING) << "valueptr is nullptr."; return {}; } std::vector cur_value = {}; - if (utils::isa(value)) { - if (!value->cast()->value().empty()) { - if (value->cast()->value().front()->type()->number_type() == kNumberTypeInt64) { - auto origin_value = GetValue>(value); - for (size_t index = 0; index < origin_value.size(); ++index) { - cur_value.push_back(static_cast(origin_value[index])); - } - } else { - cur_value = GetValue>(value); - } + if (api::utils::isa(value)) { + if (!value->cast()->value().empty()) { + auto origin_value = api::GetValue>(value); + (void)std::transform(origin_value.begin(), origin_value.end(), std::back_inserter(cur_value), + [](int64_t index) { return static_cast(index); }); } } else { - if (value->type()->number_type() == kNumberTypeInt64) { - cur_value.push_back(static_cast(GetValue(value))); - } else { - cur_value.push_back(GetValue(value)); - } + cur_value.push_back(static_cast(api::GetValue(value))); } return cur_value; } -size_t GetTupleGetItemOutIndex(const CNodePtr &tuple_get_item) { +size_t GetTupleGetItemOutIndex(const api::CNodePtr &tuple_get_item) { MS_ASSERT(tuple_get_item != nullptr); if (tuple_get_item->size() != kTupleGetItemInputSize) { MS_LOG(ERROR) << "The node tuple_get_item must have 2 inputs!"; @@ -137,7 +146,7 @@ size_t GetTupleGetItemOutIndex(const CNodePtr &tuple_get_item) { } auto output_index_value_node = tuple_get_item->input(kInputNodeOutputIndexInTupleGetItem); MS_ASSERT(output_index_value_node != nullptr); - auto value_node = output_index_value_node->cast(); + auto value_node = output_index_value_node->cast(); MS_ASSERT(value_node != nullptr); auto value_vec = CastToInt(value_node->value()); if (value_vec.empty()) { @@ -146,19 +155,19 @@ size_t GetTupleGetItemOutIndex(const CNodePtr &tuple_get_item) { } return IntToSize(value_vec.front()); } -STATUS GetOutputShapesFromCNode(const CNodePtr &cnode, std::vector *output_shapes) { - AbstractBasePtr abstract = nullptr; - if (CheckPrimitiveType(cnode, prim::kPrimTupleGetItem)) { +STATUS GetOutputShapesFromCNode(const api::CNodePtr &cnode, std::vector *output_shapes) { + api::AbstractBasePtr abstract = nullptr; + if (CheckPrimitiveType(cnode, api::MakeShared())) { auto tuple_inputs = cnode->inputs(); MS_ASSERT(tuple_inputs.size() == kTupleGetItemInputSize); auto get_item_input_cnode = tuple_inputs.at(1); MS_ASSERT(get_item_input_cnode != nullptr); auto idx = GetTupleGetItemOutIndex(cnode); - if (!utils::isa(get_item_input_cnode->abstract())) { + if (!api::utils::isa(get_item_input_cnode->abstract())) { MS_LOG(ERROR) << "TupleGetItem's abstract is not AbstractTuple"; return RET_ERROR; } - auto abstract_tuple = utils::cast(get_item_input_cnode->abstract()); + auto abstract_tuple = api::utils::cast(get_item_input_cnode->abstract()); auto abstract_list = abstract_tuple->elements(); if (abstract_list.size() <= idx) { MS_LOG(ERROR) << "AbstractTuple's size is smaller than expect"; @@ -172,8 +181,8 @@ STATUS GetOutputShapesFromCNode(const CNodePtr &cnode, std::vector MS_LOG(ERROR) << "abstract cnode is nullptr. " << cnode->fullname_with_scope(); return RET_ERROR; } - if (utils::isa(abstract)) { - auto abstract_tuple = utils::cast(abstract); + if (api::utils::isa(abstract)) { + auto abstract_tuple = api::utils::cast(abstract); auto abstract_list = abstract_tuple->elements(); for (const auto &elem : abstract_list) { ShapeVector shape_vector; @@ -203,7 +212,7 @@ STATUS GetOutputShapesFromCNode(const CNodePtr &cnode, std::vector return RET_OK; } -STATUS GetInputShapeFromCNode(const mindspore::CNodePtr &cnode, size_t input_idx, ShapeVector *shape) { +STATUS GetInputShapeFromCNode(const api::CNodePtr &cnode, size_t input_idx, ShapeVector *shape) { if (shape == nullptr) { MS_LOG(ERROR) << "shape is nullptr."; return RET_ERROR; @@ -220,7 +229,7 @@ STATUS GetInputShapeFromCNode(const mindspore::CNodePtr &cnode, size_t input_idx return RET_OK; } -STATUS FetchShapeFromAbstract(const abstract::AbstractBasePtr &abstract, ShapeVector *shape) { +STATUS FetchShapeFromAbstract(const api::AbstractBasePtr &abstract, ShapeVector *shape) { if (shape == nullptr) { MS_LOG(ERROR) << "shape is nullptr."; return RET_ERROR; @@ -229,19 +238,19 @@ STATUS FetchShapeFromAbstract(const abstract::AbstractBasePtr &abstract, ShapeVe MS_LOG(ERROR) << "abstract of cnode is invalid."; return RET_ERROR; } - if (!utils::isa(abstract)) { + if (!api::utils::isa(abstract)) { MS_LOG(ERROR) << "abstract of cnode is invalid."; return RET_ERROR; } - auto abstract_tensor = abstract->cast(); - if (!utils::isa(abstract_tensor->BuildShape())) { + auto abstract_tensor = abstract->cast(); + if (!api::utils::isa(abstract_tensor->shape())) { MS_LOG(ERROR) << "shape of cnode's output is invalid."; return RET_ERROR; } - *shape = utils::cast(abstract_tensor->BuildShape())->shape(); + *shape = api::utils::cast(abstract_tensor->shape())->shape(); return RET_OK; } -STATUS FetchTypeIdFromAbstract(const abstract::AbstractBasePtr &abstract, TypeId *type_id) { +STATUS FetchTypeIdFromAbstract(const api::AbstractBasePtr &abstract, TypeId *type_id) { if (type_id == nullptr) { MS_LOG(ERROR) << "type id is nullptr."; return RET_ERROR; @@ -250,16 +259,16 @@ STATUS FetchTypeIdFromAbstract(const abstract::AbstractBasePtr &abstract, TypeId MS_LOG(ERROR) << "abstract of cnode is invalid."; return RET_ERROR; } - if (!utils::isa(abstract)) { + if (!api::utils::isa(abstract)) { MS_LOG(ERROR) << "abstract of cnode is invalid."; return RET_ERROR; } - auto abstract_tensor = abstract->cast(); + auto abstract_tensor = abstract->cast(); if (abstract_tensor->element() == nullptr) { MS_LOG(ERROR) << "element of abstract_tensor is nullptr."; return RET_ERROR; } - auto type_ptr = abstract_tensor->element()->GetTypeTrack(); + auto type_ptr = abstract_tensor->element()->type(); if (type_ptr == nullptr) { MS_LOG(ERROR) << "type_ptr of abstract_tensor is nullptr."; return RET_ERROR; @@ -268,18 +277,18 @@ STATUS FetchTypeIdFromAbstract(const abstract::AbstractBasePtr &abstract, TypeId return RET_OK; } -int GetAnfNodeOutputShape(const AnfNodePtr &input, ShapeVector *shape_vector) { +int GetAnfNodeOutputShape(const api::AnfNodePtr &input, ShapeVector *shape_vector) { if (shape_vector == nullptr) { MS_LOG(ERROR) << "shape vector is nullptr." << input->fullname_with_scope(); return RET_ERROR; } - if (utils::isa(input)) { + if (api::utils::isa(input)) { if (GetShapeVectorFromParameter(input, shape_vector) != RET_OK) { MS_LOG(ERROR) << "get output shape for preprocessor failed. " << input->fullname_with_scope(); return RET_ERROR; } - } else if (utils::isa(input)) { - auto input_cnode = input->cast(); + } else if (api::utils::isa(input)) { + auto input_cnode = input->cast(); std::vector output_shapes; if (GetOutputShapesFromCNode(input_cnode, &output_shapes) != RET_OK) { MS_LOG(ERROR) << "get output shapes from cnode failed. " << input_cnode->fullname_with_scope(); @@ -318,27 +327,27 @@ std::string TypeIdToString(TypeId type_id) { return type_str; } -bool CheckInputs(const CNodePtr &cnode) { +bool CheckInputs(const api::CNodePtr &cnode) { if (cnode == nullptr) { MS_LOG(ERROR) << "cnode is nullptr."; return false; } - if (std::any_of(cnode->inputs().begin(), cnode->inputs().end(), - [](const AnfNodePtr &anf_node) { return anf_node == nullptr; })) { + auto inputs = cnode->inputs(); + if (std::any_of(inputs.begin(), inputs.end(), [](const api::AnfNodePtr &anf_node) { return anf_node == nullptr; })) { MS_LOG(ERROR) << "input is nullptr."; return false; } return true; } -std::string GetCustomOutputName(const AnfNodePtr &node) { +std::string GetCustomOutputName(const api::AnfNodePtr &node) { std::string output_name; - auto input_cnode = node->cast(); + auto input_cnode = node->cast(); if (input_cnode == nullptr) { MS_LOG(ERROR) << "custom node should be cnode. " << node->fullname_with_scope(); return ""; } if (input_cnode->GetAttr(kOutputsNames) != nullptr) { - auto output_names = GetValue>(input_cnode->GetAttr(kOutputsNames)); + auto output_names = api::GetValue>(input_cnode->GetAttr(kOutputsNames)); if (output_names.size() == 1) { output_name = output_names.at(0); } else { @@ -349,19 +358,19 @@ std::string GetCustomOutputName(const AnfNodePtr &node) { } return output_name; } -tensor::TensorPtr CreateTensorInfo(const void *data, size_t data_size, const std::vector &shape, - TypeId data_type) { - tensor::TensorPtr tensor_info = nullptr; - if (shape.empty() && data_size == abstract::TypeIdSize(data_type)) { +api::TensorPtr CreateTensorInfo(const void *data, size_t data_size, const std::vector &shape, + TypeId data_type) { + api::TensorPtr tensor_info = nullptr; + if (shape.empty() && data_size == TypeIdSize(data_type)) { ShapeVector scalar_shape = {1}; - tensor_info = std::make_shared(data_type, scalar_shape); + tensor_info = api::MakeShared(data_type, scalar_shape); if (tensor_info == nullptr) { MS_LOG(ERROR) << "new tensor init failed"; return nullptr; } tensor_info->set_shape({}); } else { - tensor_info = std::make_shared(data_type, shape); + tensor_info = api::MakeShared(data_type, shape); if (tensor_info == nullptr) { MS_LOG(ERROR) << "new tensor init failed"; return nullptr; @@ -374,7 +383,7 @@ tensor::TensorPtr CreateTensorInfo(const void *data, size_t data_size, const std MS_LOG(ERROR) << "input tensor data is nullptr"; return nullptr; } - auto ret = memcpy_s(tensor_info->data_c(), tensor_info->data().nbytes(), data, data_size); + auto ret = memcpy_s(tensor_info->data(), tensor_info->Size(), data, data_size); if (ret != EOK) { MS_LOG(ERROR) << "memcpy_s error : " << ret; return nullptr; @@ -382,7 +391,7 @@ tensor::TensorPtr CreateTensorInfo(const void *data, size_t data_size, const std return tensor_info; } -AbstractBasePtr CreateTensorAbstract(const std::vector &shape, TypeId data_type) { +api::AbstractBasePtr CreateTensorAbstract(const std::vector &shape, TypeId data_type) { auto tensor_info = dpico::CreateTensorInfo(nullptr, 0, shape, data_type); if (tensor_info == nullptr) { MS_LOG(ERROR) << "Create tensor info failed"; @@ -396,7 +405,7 @@ AbstractBasePtr CreateTensorAbstract(const std::vector &shape, TypeId d return abstract; } -int InitParameterFromTensorInfo(const ParameterPtr ¶m_node, const tensor::TensorPtr &tensor_info) { +int InitParameterFromTensorInfo(const api::ParameterPtr ¶m_node, const api::TensorPtr &tensor_info) { if (tensor_info == nullptr) { MS_LOG(ERROR) << "tensor info is nullptr."; return RET_ERROR; @@ -411,7 +420,7 @@ int InitParameterFromTensorInfo(const ParameterPtr ¶m_node, const tensor::Te return RET_OK; } -abstract::AbstractBasePtr GetCNodeInputAbstract(const CNodePtr &cnode, size_t index) { +api::AbstractBasePtr GetCNodeInputAbstract(const api::CNodePtr &cnode, size_t index) { if (cnode == nullptr) { MS_LOG(ERROR) << "CNodePtr is nullptr"; return nullptr; @@ -427,23 +436,23 @@ abstract::AbstractBasePtr GetCNodeInputAbstract(const CNodePtr &cnode, size_t in return nullptr; } - abstract::AbstractBasePtr abstract = nullptr; - if (utils::isa(input)) { - auto parameter = input->cast(); + api::AbstractBasePtr abstract = nullptr; + if (api::utils::isa(input)) { + auto parameter = input->cast(); abstract = parameter->abstract(); - } else if (utils::isa(input)) { - auto input_cnode = input->cast(); - if (CheckPrimitiveType(input_cnode, prim::kPrimTupleGetItem)) { + } else if (api::utils::isa(input)) { + auto input_cnode = input->cast(); + if (CheckPrimitiveType(input_cnode, api::MakeShared())) { auto tuple_inputs = input_cnode->inputs(); MS_ASSERT(tuple_inputs.size() == kTupleGetItemInputSize); auto get_item_input_cnode = tuple_inputs.at(1); MS_ASSERT(get_item_input_cnode != nullptr); auto idx = GetTupleGetItemOutIndex(input_cnode); - if (!utils::isa(get_item_input_cnode->abstract())) { + if (!api::utils::isa(get_item_input_cnode->abstract())) { MS_LOG(ERROR) << "TupleGetItem's abstract is not AbstractTuple"; return nullptr; } - auto abstract_tuple = utils::cast(get_item_input_cnode->abstract()); + auto abstract_tuple = api::utils::cast(get_item_input_cnode->abstract()); auto abstract_list = abstract_tuple->elements(); if (abstract_list.size() <= idx) { MS_LOG(ERROR) << "AbstractTuple's size is smaller than expect"; @@ -460,24 +469,24 @@ abstract::AbstractBasePtr GetCNodeInputAbstract(const CNodePtr &cnode, size_t in return abstract; } -abstract::AbstractBasePtr GetAbstractFromAnfNode(const AnfNodePtr &node) { - AbstractBasePtr abstract = nullptr; - if (utils::isa(node)) { - auto parameter = node->cast(); +api::AbstractBasePtr GetAbstractFromAnfNode(const api::AnfNodePtr &node) { + api::AbstractBasePtr abstract = nullptr; + if (api::utils::isa(node)) { + auto parameter = node->cast(); abstract = parameter->abstract(); - } else if (utils::isa(node)) { - auto cnode = node->cast(); - if (CheckPrimitiveType(cnode, prim::kPrimTupleGetItem)) { + } else if (api::utils::isa(node)) { + auto cnode = node->cast(); + if (CheckPrimitiveType(cnode, api::MakeShared())) { auto tuple_inputs = cnode->inputs(); MS_ASSERT(tuple_inputs.size() == kTupleGetItemInputSize); auto get_item_input_cnode = tuple_inputs.at(1); MS_ASSERT(get_item_input_cnode != nullptr); auto idx = GetTupleGetItemOutIndex(cnode); - if (!utils::isa(get_item_input_cnode->abstract())) { + if (!api::utils::isa(get_item_input_cnode->abstract())) { MS_LOG(ERROR) << "TupleGetItem's abstract is not AbstractTuple"; return nullptr; } - auto abstract_tuple = utils::cast(get_item_input_cnode->abstract()); + auto abstract_tuple = api::utils::cast(get_item_input_cnode->abstract()); auto abstract_list = abstract_tuple->elements(); if (abstract_list.size() <= idx) { MS_LOG(ERROR) << "AbstractTuple's size is smaller than expect"; @@ -491,8 +500,8 @@ abstract::AbstractBasePtr GetAbstractFromAnfNode(const AnfNodePtr &node) { return abstract; } -ParameterPtr BuildIntValueParameterNode(const api::FuncGraphPtr &func_graph, const int32_t &data, - const std::string &node_name) { +api::ParameterPtr BuildIntValueParameterNode(const api::FuncGraphPtr &func_graph, const int32_t &data, + const std::string &node_name) { MS_ASSERT(func_graph != nullptr); auto param_node = func_graph->add_parameter(); param_node->set_name(node_name); @@ -511,8 +520,8 @@ ParameterPtr BuildIntValueParameterNode(const api::FuncGraphPtr &func_graph, con return param_node; } -ParameterPtr BuildIntVecParameterNode(const api::FuncGraphPtr &func_graph, const std::vector &data, - const std::string &node_name) { +api::ParameterPtr BuildIntVecParameterNode(const api::FuncGraphPtr &func_graph, const std::vector &data, + const std::string &node_name) { MS_ASSERT(func_graph != nullptr); MS_ASSERT(data.size() != 0); auto param_node = func_graph->add_parameter(); @@ -534,8 +543,9 @@ ParameterPtr BuildIntVecParameterNode(const api::FuncGraphPtr &func_graph, const return param_node; } -ParameterPtr BuildIntVec2DParameterNode(const api::FuncGraphPtr &func_graph, - const std::vector> &data, const std::string &node_name) { +api::ParameterPtr BuildIntVec2DParameterNode(const api::FuncGraphPtr &func_graph, + const std::vector> &data, + const std::string &node_name) { MS_ASSERT(func_graph != nullptr); MS_ASSERT(data.size() != 0); auto param_node = func_graph->add_parameter(); @@ -564,8 +574,8 @@ ParameterPtr BuildIntVec2DParameterNode(const api::FuncGraphPtr &func_graph, return param_node; } -ParameterPtr BuildFloatValueParameterNode(const api::FuncGraphPtr &func_graph, const float &data, - const std::string &node_name) { +api::ParameterPtr BuildFloatValueParameterNode(const api::FuncGraphPtr &func_graph, const float &data, + const std::string &node_name) { MS_ASSERT(func_graph != nullptr); auto param_node = func_graph->add_parameter(); param_node->set_name(node_name); @@ -583,15 +593,15 @@ ParameterPtr BuildFloatValueParameterNode(const api::FuncGraphPtr &func_graph, c return param_node; } -CNodePtr GenTransposeNode(const api::FuncGraphPtr &func_graph, const AnfNodePtr &input_node, - const std::vector &perm, const std::string &cnode_name) { +api::CNodePtr GenTransposeNode(const api::FuncGraphPtr &func_graph, const api::AnfNodePtr &input_node, + const std::vector &perm, const std::string &cnode_name) { MS_ASSERT(func_graph != nullptr && input_node != nullptr); auto perm_node = BuildIntVecParameterNode(func_graph, perm, cnode_name + "_perm"); if (perm_node == nullptr) { MS_LOG(ERROR) << "new perm_node error"; return nullptr; } - auto trans_prim = std::make_shared(); + auto trans_prim = api::MakeShared(); if (trans_prim == nullptr) { MS_LOG(ERROR) << "new trans_prim failed"; return nullptr; @@ -612,12 +622,12 @@ CNodePtr GenTransposeNode(const api::FuncGraphPtr &func_graph, const AnfNodePtr return cnode; } -tensor::TensorPtr GetTensorInfo(const AnfNodePtr &node) { +api::TensorPtr GetTensorInfo(const api::AnfNodePtr &node) { MS_ASSERT(node != nullptr); - if (!utils::isa(node)) { - if (utils::isa(node)) { - auto valueNode = node->cast(); - auto value = std::dynamic_pointer_cast(valueNode->value()); + if (!api::utils::isa(node)) { + if (api::utils::isa(node)) { + auto valueNode = node->cast(); + auto value = valueNode->value()->cast(); if (value != nullptr) { return value; } @@ -625,53 +635,42 @@ tensor::TensorPtr GetTensorInfo(const AnfNodePtr &node) { MS_LOG(DEBUG) << "get lite param value node neither parameter node or value node"; return nullptr; } - auto param = node->cast(); + auto param = node->cast(); if (param == nullptr) { MS_LOG(ERROR) << "param is nullptr."; return nullptr; } - auto tensor_info = std::dynamic_pointer_cast(param->default_param()); + auto tensor_info = param->default_param()->cast(); return tensor_info; } -std::vector> CastToVec2DInt(const ValuePtr &value) { +std::vector> CastToVec2DInt(const api::ValuePtr &value) { if (value == nullptr) { MS_LOG(WARNING) << "valueptr is nullptr."; return {}; } std::vector> result_value; - if (utils::isa(value)) { - if (value->cast() - ->value() - .front() - ->cast() - ->value() - .front() - ->type() - ->number_type() == kNumberTypeInt64) { - auto origin_value = GetValue>>(value); - for (size_t i = 0; i < origin_value.size(); ++i) { - std::vector cur_value; - for (size_t j = 0; j < origin_value.at(i).size(); ++j) { - cur_value.push_back(static_cast(origin_value[i][j])); - } - result_value.push_back(cur_value); + if (api::utils::isa(value)) { + auto origin_value = api::GetValue>>(value); + for (auto &vec : origin_value) { + std::vector cur_value; + for (size_t j = 0; j < vec.size(); ++j) { + cur_value.push_back(static_cast(vec[j])); } - } else { - result_value = GetValue>>(value); + result_value.push_back(cur_value); } } return result_value; } -bool GetBoolAttr(const AnfNodePtr &node, const std::string &attr_name) { - auto cnode = node->cast(); +bool GetBoolAttr(const api::AnfNodePtr &node, const std::string &attr_name) { + auto cnode = node->cast(); if (cnode == nullptr) { MS_LOG(ERROR) << "cur node is not a cnode. " << node->fullname_with_scope(); return false; } - auto primitive = GetValueNode(cnode->input(0)); + auto primitive = api::GetValueNode(cnode->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr:" << cnode->fullname_with_scope(); return false; @@ -681,10 +680,10 @@ bool GetBoolAttr(const AnfNodePtr &node, const std::string &attr_name) { MS_LOG(ERROR) << "There is no attr named " << attr_name << " for node " << cnode->fullname_with_scope(); return false; } - return GetValue(value_ptr); + return api::GetValue(value_ptr); } -STATUS GetDataTypeAndShape(const ParameterPtr ¶m_node, TypeId *data_type, ShapeVector *shape_vector) { +STATUS GetDataTypeAndShape(const api::ParameterPtr ¶m_node, TypeId *data_type, ShapeVector *shape_vector) { if (param_node == nullptr) { MS_LOG(ERROR) << "param node is nullptr."; return RET_ERROR; @@ -702,24 +701,24 @@ STATUS GetDataTypeAndShape(const ParameterPtr ¶m_node, TypeId *data_type, Sh MS_LOG(ERROR) << "Abstract of parameter is nullptr, " << param_node->name(); return RET_ERROR; } - if (!utils::isa(abstract_base)) { + if (!api::utils::isa(abstract_base)) { MS_LOG(ERROR) << "Abstract of parameter should be abstract tensor, " << param_node->name(); return RET_ERROR; } - auto abstract_tensor = utils::cast(abstract_base); + auto abstract_tensor = api::utils::cast(abstract_base); MS_CHECK_TRUE_MSG(abstract_tensor != nullptr, RET_ERROR, "Cast to abstract tensor failed!"); - auto typePtr = abstract_tensor->element()->GetTypeTrack(); + auto typePtr = abstract_tensor->element()->type(); MS_ASSERT(typePtr != nullptr); *data_type = typePtr->type_id(); - if (!utils::isa(abstract_tensor->BuildShape())) { + if (!api::utils::isa(abstract_tensor->shape())) { MS_LOG(ERROR) << "Shape of Abstract of parameter should be ShapePtr, " << param_node->name(); return RET_ERROR; } - *shape_vector = utils::cast(abstract_tensor->BuildShape())->shape(); + *shape_vector = api::utils::cast(abstract_tensor->shape())->shape(); return RET_OK; } -STATUS GetShapeVectorFromStringTensor(const tensor::TensorPtr &tensor_info, ShapeVector *shape_vector, size_t *offset) { +STATUS GetShapeVectorFromStringTensor(const api::TensorPtr &tensor_info, ShapeVector *shape_vector, size_t *offset) { if (tensor_info == nullptr) { MS_LOG(ERROR) << "tensor info is nullptr."; return RET_ERROR; @@ -738,7 +737,7 @@ STATUS GetShapeVectorFromStringTensor(const tensor::TensorPtr &tensor_info, Shap return RET_ERROR; } shape_vector->clear(); - auto tensor_data = reinterpret_cast(tensor_info->data_c()); + auto tensor_data = reinterpret_cast(tensor_info->data()); std::string shape_str; std::string shape_size_str; *offset = 0; diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.h b/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.h index 5a95d8fa0a..cf45196a1a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/anf_util.h @@ -19,11 +19,11 @@ #include #include +#include "mindapi/ir/tensor.h" #include "include/errorcode.h" -#include "utils/log_adapter.h" -#include "ir/anf.h" -#include "api/ir/func_graph.h" -#include "ops/primitive_c.h" +#include "mindapi/base/logging.h" +#include "mindapi/ir/anf.h" +#include "mindapi/ir/func_graph.h" using mindspore::lite::RET_ERROR; using mindspore::lite::RET_NO_CHANGE; @@ -32,44 +32,41 @@ using mindspore::lite::STATUS; namespace mindspore { namespace dpico { -bool CheckPrimitiveType(const mindspore::AnfNodePtr &node, const mindspore::PrimitivePtr &primitive_type); -STATUS GetPrimitiveType(const mindspore::AnfNodePtr &node, std::string *name); -STATUS GetShapeVectorFromParameter(const mindspore::AnfNodePtr &weight, ShapeVector *shape_vector); -std::vector CastToInt(const mindspore::ValuePtr &value); -size_t GetTupleGetItemOutIndex(const mindspore::CNodePtr &tuple_get_item); -STATUS GetOutputShapesFromCNode(const mindspore::CNodePtr &cnode, std::vector *output_shapes); -STATUS GetInputShapeFromCNode(const mindspore::CNodePtr &cnode, size_t input_idx, ShapeVector *shape); -STATUS FetchShapeFromAbstract(const mindspore::abstract::AbstractBasePtr &abstract, ShapeVector *shape); -STATUS FetchTypeIdFromAbstract(const mindspore::abstract::AbstractBasePtr &abstract, TypeId *type_id); -int GetAnfNodeOutputShape(const AnfNodePtr &input, ShapeVector *shape_vector); +bool CheckPrimitiveType(const api::AnfNodePtr &node, const api::PrimitivePtr &primitive_type); +STATUS GetPrimitiveType(const api::AnfNodePtr &node, std::string *name); +STATUS GetShapeVectorFromParameter(const api::AnfNodePtr &weight, ShapeVector *shape_vector); +std::vector CastToInt(const api::ValuePtr &value); +size_t GetTupleGetItemOutIndex(const api::CNodePtr &tuple_get_item); +STATUS GetOutputShapesFromCNode(const api::CNodePtr &cnode, std::vector *output_shapes); +STATUS GetInputShapeFromCNode(const api::CNodePtr &cnode, size_t input_idx, ShapeVector *shape); +STATUS FetchShapeFromAbstract(const api::AbstractBasePtr &abstract, ShapeVector *shape); +STATUS FetchTypeIdFromAbstract(const api::AbstractBasePtr &abstract, TypeId *type_id); +int GetAnfNodeOutputShape(const api::AnfNodePtr &input, ShapeVector *shape_vector); std::string TypeIdToString(TypeId type_id); -bool CheckInputs(const mindspore::CNodePtr &cnode); -std::string GetCustomOutputName(const AnfNodePtr &node); -mindspore::tensor::TensorPtr CreateTensorInfo(const void *data, size_t data_size, const std::vector &shape, - mindspore::TypeId data_type); -mindspore::AbstractBasePtr CreateTensorAbstract(const std::vector &shape, mindspore::TypeId data_type); -int InitParameterFromTensorInfo(const mindspore::ParameterPtr ¶m_node, - const mindspore::tensor::TensorPtr &tensor_info); -mindspore::abstract::AbstractBasePtr GetCNodeInputAbstract(const mindspore::CNodePtr &cnode, size_t index); -mindspore::abstract::AbstractBasePtr GetAbstractFromAnfNode(const AnfNodePtr &cnode); -mindspore::ParameterPtr BuildIntValueParameterNode(const api::FuncGraphPtr &func_graph, const int32_t &data, - const std::string &node_name); -mindspore::ParameterPtr BuildIntVecParameterNode(const api::FuncGraphPtr &func_graph, const std::vector &data, - const std::string &node_name); -mindspore::ParameterPtr BuildIntVec2DParameterNode(const api::FuncGraphPtr &func_graph, - const std::vector> &data, - const std::string &node_name); -mindspore::ParameterPtr BuildFloatValueParameterNode(const api::FuncGraphPtr &func_graph, const float &data, - const std::string &node_name); -mindspore::CNodePtr GenTransposeNode(const api::FuncGraphPtr &func_graph, const mindspore::AnfNodePtr &input_node, - const std::vector &perm, const std::string &cnode_name); -mindspore::tensor::TensorPtr GetTensorInfo(const mindspore::AnfNodePtr &node); -std::vector> CastToVec2DInt(const mindspore::ValuePtr &value); -bool GetBoolAttr(const mindspore::AnfNodePtr &node, const std::string &attr_name); -STATUS GetDataTypeAndShape(const mindspore::ParameterPtr ¶m_node, mindspore::TypeId *data_type, - ShapeVector *shape_vector); -STATUS GetShapeVectorFromStringTensor(const mindspore::tensor::TensorPtr &tensor_info, ShapeVector *shape_vector, - size_t *offset); +bool CheckInputs(const api::CNodePtr &cnode); +std::string GetCustomOutputName(const api::AnfNodePtr &node); +api::TensorPtr CreateTensorInfo(const void *data, size_t data_size, const std::vector &shape, + TypeId data_type); +api::AbstractBasePtr CreateTensorAbstract(const std::vector &shape, TypeId data_type); +int InitParameterFromTensorInfo(const api::ParameterPtr ¶m_node, const api::TensorPtr &tensor_info); +api::AbstractBasePtr GetCNodeInputAbstract(const api::CNodePtr &cnode, size_t index); +api::AbstractBasePtr GetAbstractFromAnfNode(const api::AnfNodePtr &cnode); +api::ParameterPtr BuildIntValueParameterNode(const api::FuncGraphPtr &func_graph, const int32_t &data, + const std::string &node_name); +api::ParameterPtr BuildIntVecParameterNode(const api::FuncGraphPtr &func_graph, const std::vector &data, + const std::string &node_name); +api::ParameterPtr BuildIntVec2DParameterNode(const api::FuncGraphPtr &func_graph, + const std::vector> &data, + const std::string &node_name); +api::ParameterPtr BuildFloatValueParameterNode(const api::FuncGraphPtr &func_graph, const float &data, + const std::string &node_name); +api::CNodePtr GenTransposeNode(const api::FuncGraphPtr &func_graph, const api::AnfNodePtr &input_node, + const std::vector &perm, const std::string &cnode_name); +api::TensorPtr GetTensorInfo(const api::AnfNodePtr &node); +std::vector> CastToVec2DInt(const api::ValuePtr &value); +bool GetBoolAttr(const api::AnfNodePtr &node, const std::string &attr_name); +STATUS GetDataTypeAndShape(const api::ParameterPtr ¶m_node, TypeId *data_type, ShapeVector *shape_vector); +STATUS GetShapeVectorFromStringTensor(const api::TensorPtr &tensor_info, ShapeVector *shape_vector, size_t *offset); inline size_t IntToSize(int u) { if (u < 0) { MS_LOG(WARNING) << "The int value(" << u << ") is less than 0."; diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/check_base.h b/mindspore/lite/tools/converter/adapter/dpico/common/check_base.h index f08bc0e59f..9a2d30165f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/check_base.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/check_base.h @@ -32,6 +32,12 @@ #define kNCHW_C 1 #define kNCHW_H 2 #define kNCHW_W 3 +#ifdef Debug +#include +#define MS_ASSERT(f) assert(f) +#else +#define MS_ASSERT(f) ((void)0) +#endif #define SIZE_MUL_OVERFLOW(x, y) (((x) == 0) ? false : (SIZE_MAX / (x)) < (y)) #define INT_MUL_OVERFLOW(x, y) \ diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.cc b/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.cc index 80a2e6d890..342fe5a8ba 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.cc @@ -14,12 +14,15 @@ * limitations under the License. */ #include "common/data_transpose_utils.h" -#include #include #include #include #include +#include +#include "third_party/securec/include/securec.h" #include "common/check_base.h" +#include "common/float16.h" +#include "mindapi/base/logging.h" namespace mindspore { namespace dpico { @@ -51,7 +54,7 @@ void MoveData(float *matrix, int idx, int row, int col) { // idx is from the ne } // namespace int DeduceDimConvertion(mindspore::Format src_format, mindspore::Format dst_format, std::vector *perm) { - MS_ASSERT(perm != nullptr); + MS_CHECK_TRUE_MSG(perm != nullptr, RET_ERROR, "perm is nullptr."); if (kTensorFormatMap.find(src_format) == kTensorFormatMap.end() || kTensorFormatMap.find(dst_format) == kTensorFormatMap.end()) { MS_LOG(ERROR) << "src_format or dst_format is error."; @@ -121,9 +124,9 @@ STATUS TransposeData(const ShapeVector &origin_shape, const ShapeVector &cur_sha } template -STATUS DoTransposeData(const tensor::TensorPtr &tensor, mindspore::Format src_format, mindspore::Format dst_format) { +STATUS DoTransposeData(const api::TensorPtr &tensor, mindspore::Format src_format, mindspore::Format dst_format) { MS_ASSERT(tensor != nullptr); - auto origin_shape = tensor->shape_c(); + auto origin_shape = tensor->shape(); if (origin_shape.size() != kDims4) { MS_LOG(ERROR) << "Filter dim-num is not supported, dim-num: " << origin_shape.size(); return lite::RET_ERROR; @@ -159,7 +162,7 @@ STATUS DoTransposeData(const tensor::TensorPtr &tensor, mindspore::Format src_fo } std::vector buf(count); - auto origin_weight_data = tensor->data_c(); + auto origin_weight_data = tensor->data(); if (origin_weight_data == nullptr) { MS_LOG(ERROR) << "origin_weight_data is nullptr."; return RET_ERROR; @@ -173,7 +176,7 @@ STATUS DoTransposeData(const tensor::TensorPtr &tensor, mindspore::Format src_fo MS_LOG(ERROR) << "tensor size shouldn't be 0"; return RET_ERROR; } - if (memcpy_s(tensor->data_c(), tensor->Size(), buf.data(), count * sizeof(T)) != EOK) { + if (memcpy_s(tensor->data(), tensor->Size(), buf.data(), count * sizeof(T)) != EOK) { MS_LOG(ERROR) << "memcpy_s failed."; return RET_ERROR; } @@ -181,12 +184,12 @@ STATUS DoTransposeData(const tensor::TensorPtr &tensor, mindspore::Format src_fo return RET_OK; } -STATUS TransFilterFormat(const tensor::TensorPtr &tensor, mindspore::Format src_format, mindspore::Format dst_format) { +STATUS TransFilterFormat(const api::TensorPtr &tensor, mindspore::Format src_format, mindspore::Format dst_format) { if (tensor == nullptr) { MS_LOG(ERROR) << "tensor is nullptr."; return RET_ERROR; } - std::unordered_map> + std::unordered_map> trans_func = {{kNumberTypeFloat32, DoTransposeData}, {kNumberTypeUInt8, DoTransposeData}, {kNumberTypeInt8, DoTransposeData}, diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.h b/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.h index 675320a4cd..5845c515f8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/data_transpose_utils.h @@ -18,9 +18,9 @@ #define DPICO_COMMON_DATA_TRANSPOSE_UTILS_H #include -#include "include/api/format.h" -#include "ir/tensor.h" -#include "utils/log_adapter.h" +#include "mindapi/base/format.h" +#include "mindapi/ir/tensor.h" +#include "mindapi/base/logging.h" #include "include/errorcode.h" #include "common/op_enum.h" using mindspore::lite::RET_ERROR; @@ -79,7 +79,7 @@ STATUS NCHW2NHWC(T *src_data, T *dst_data, std::vector shape) { return RET_OK; } -STATUS TransFilterFormat(const mindspore::tensor::TensorPtr &tensor, mindspore::Format src_format, +STATUS TransFilterFormat(const mindspore::api::TensorPtr &tensor, mindspore::Format src_format, mindspore::Format dst_format); void TransposeMatrix(float *matrix, int row, int col); diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/dynamic_library_loader.cc b/mindspore/lite/tools/converter/adapter/dpico/common/dynamic_library_loader.cc deleted file mode 100644 index 7f3e9fd347..0000000000 --- a/mindspore/lite/tools/converter/adapter/dpico/common/dynamic_library_loader.cc +++ /dev/null @@ -1,89 +0,0 @@ -/** - * 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 "common/dynamic_library_loader.h" -#include -#ifndef _WIN32 -#include -#else -#include -#undef ERROR -#undef SM_DEBUG -#endif -#include "include/errorcode.h" -#include "common/file_util.h" - -namespace mindspore { -namespace dpico { -int DynamicLibraryLoader::Open(const std::string &lib_path) { - if (handler_ != nullptr) { - return RET_ERROR; - } - std::string real_path = RealPath(lib_path.c_str()); - -#ifndef _WIN32 -#ifndef ENABLE_ARM - handler_ = dlopen(real_path.c_str(), RTLD_LAZY | RTLD_DEEPBIND); -#else - handler_ = dlopen(real_path.c_str(), RTLD_LAZY); -#endif -#else - handler_ = LoadLibrary(real_path.c_str()); -#endif - if (handler_ == nullptr) { - MS_LOG(ERROR) << "handler is nullptr."; - return RET_ERROR; - } - return RET_OK; -} - -void *DynamicLibraryLoader::GetFunc(const std::string &func_name) { -#ifndef _WIN32 - return dlsym(handler_, func_name.c_str()); -#else - auto func = GetProcAddress(reinterpret_cast(handler_), func_name.c_str()); - return reinterpret_cast(func); -#endif -} - -int DynamicLibraryLoader::Close() { - if (handler_ == nullptr) { - return RET_OK; - } -#ifndef _WIN32 - auto close_res = dlclose(handler_); - if (close_res != 0) { - MS_LOG(ERROR) << "can not close handler"; - return RET_ERROR; - } -#else - auto close_res = FreeLibrary(reinterpret_cast(handler_)); - if (close_res == 0) { - MS_LOG(ERROR) << "can not close handler"; - return RET_ERROR; - } -#endif - handler_ = nullptr; - return RET_OK; -} - -DynamicLibraryLoader::~DynamicLibraryLoader() { - if (handler_ != nullptr) { - Close(); - } -} -} // namespace dpico -} // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/dynamic_library_loader.h b/mindspore/lite/tools/converter/adapter/dpico/common/dynamic_library_loader.h deleted file mode 100644 index 4e87db119c..0000000000 --- a/mindspore/lite/tools/converter/adapter/dpico/common/dynamic_library_loader.h +++ /dev/null @@ -1,38 +0,0 @@ -/** - * 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. - */ - -#ifndef DPICO_COMMON_DYNAMIC_LIBRARY_LOADER_H_ -#define DPICO_COMMON_DYNAMIC_LIBRARY_LOADER_H_ - -#include - -namespace mindspore { -namespace dpico { -class DynamicLibraryLoader { - public: - DynamicLibraryLoader() = default; - ~DynamicLibraryLoader(); - int Open(const std::string &lib_path); - void *GetFunc(const std::string &func_name); - int Close(); - - private: - void *handler_ = nullptr; -}; -} // namespace dpico -} // namespace mindspore - -#endif // DPICO_COMMON_DYNAMIC_LIBRARY_LOADER_H_ diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.cc b/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.cc index 6905213fef..32e898c627 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.cc @@ -15,13 +15,11 @@ */ #include "common/fetch_content.h" -#include #include -#include #include -#include #include "common/anf_util.h" #include "common/check_base.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace dpico { @@ -29,7 +27,7 @@ namespace { constexpr size_t kTensorListMinSize = 3 * sizeof(int32_t); } // namespace -int FetchFromDefaultParam(const ParameterPtr ¶m_node, DataInfo *data_info) { +int FetchFromDefaultParam(const api::ParameterPtr ¶m_node, DataInfo *data_info) { MS_ASSERT(param_node != nullptr && data_info != nullptr); ShapeVector shape_vector; TypeId data_type; @@ -39,7 +37,7 @@ int FetchFromDefaultParam(const ParameterPtr ¶m_node, DataInfo *data_info) { return RET_ERROR; } data_info->data_type_ = data_type; - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); size_t offset = 0; if (!shape_vector.empty() && data_type == kObjectTypeString) { status = GetShapeVectorFromStringTensor(tensor_info, &shape_vector, &offset); @@ -58,7 +56,7 @@ int FetchFromDefaultParam(const ParameterPtr ¶m_node, DataInfo *data_info) { } data_info->data_.resize(tensor_info->Size() - offset); if (EOK != memcpy_s(data_info->data_.data(), data_info->data_.size(), - static_cast(tensor_info->data_c()) + offset, tensor_info->Size() - offset)) { + static_cast(tensor_info->data()) + offset, tensor_info->Size() - offset)) { MS_LOG(ERROR) << "memcpy_s failed."; return RET_ERROR; } @@ -68,13 +66,13 @@ int FetchFromDefaultParam(const ParameterPtr ¶m_node, DataInfo *data_info) { return RET_OK; } -int FetchDataFromParameterNode(const CNodePtr &cnode, size_t index, DataInfo *data_info) { +int FetchDataFromParameterNode(const api::CNodePtr &cnode, size_t index, DataInfo *data_info) { MS_ASSERT(cnode != nullptr && data_info != nullptr); if (index >= cnode->inputs().size()) { MS_LOG(ERROR) << "input index: " << index << " is greater than cnode inputs size " << cnode->inputs().size(); return RET_ERROR; } - auto param_node = cnode->input(index)->cast(); + auto param_node = cnode->input(index)->cast(); if (param_node == nullptr) { MS_LOG(ERROR) << "input node is not parameter node."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.h b/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.h index bcc85d2c14..3857753e00 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/fetch_content.h @@ -14,13 +14,13 @@ * limitations under the License. */ -#ifndef MINDSPORE_LITE_TOOLS_ANF_EXPORTER_FETCH_CONTENT_H_ -#define MINDSPORE_LITE_TOOLS_ANF_EXPORTER_FETCH_CONTENT_H_ +#ifndef DPICO_COMMON_FETCH_CONTENT_H_ +#define DPICO_COMMON_FETCH_CONTENT_H_ #include #include -#include "ir/primitive.h" -#include "ir/func_graph.h" +#include "mindapi/ir/primitive.h" +#include "mindapi/ir/func_graph.h" namespace mindspore { namespace dpico { @@ -31,10 +31,11 @@ struct DataInfo { DataInfo() : data_type_(0) {} }; -int FetchFromDefaultParam(const ParameterPtr ¶m_node, DataInfo *data_info); +int FetchFromDefaultParam(const api::ParameterPtr ¶m_node, DataInfo *data_info); + +int FetchDataFromParameterNode(const api::CNodePtr &cnode, size_t index, DataInfo *data_info); -int FetchDataFromParameterNode(const CNodePtr &cnode, size_t index, DataInfo *data_info); int GetDataSizeFromTensor(DataInfo *data_info, int *data_size); } // namespace dpico } // namespace mindspore -#endif // MINDSPORE_LITE_TOOLS_ANF_EXPORTER_FETCH_CONTENT_H_ +#endif // DPICO_COMMON_FETCH_CONTENT_H_ diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/file_util.cc b/mindspore/lite/tools/converter/adapter/dpico/common/file_util.cc index 04d71c3836..cf3a37ef92 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/file_util.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/file_util.cc @@ -14,13 +14,14 @@ * limitations under the License. */ +#include "common/file_util.h" #include #include #include #include #include -#include "utils/log_adapter.h" -#include "common/file_util.h" +#include +#include "mindapi/base/logging.h" namespace mindspore { namespace dpico { @@ -69,7 +70,7 @@ std::string RealPath(const char *path) { MS_LOG(ERROR) << "path is nullptr"; return ""; } - if ((strlen(path)) >= PATH_MAX) { + if ((std::strlen(path)) >= PATH_MAX) { MS_LOG(ERROR) << "path is too long"; return ""; } @@ -152,7 +153,7 @@ int RemoveDir(const std::string &path) { struct dirent *dt = nullptr; dt = readdir(d); while (dt != nullptr) { - if (strcmp(dt->d_name, "..") != 0 && strcmp(dt->d_name, ".") != 0) { + if (std::strcmp(dt->d_name, "..") != 0 && std::strcmp(dt->d_name, ".") != 0) { struct stat st {}; auto file_name = str_path + std::string(dt->d_name); stat(file_name.c_str(), &st); diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/file_util.h b/mindspore/lite/tools/converter/adapter/dpico/common/file_util.h index 4eab1b24b7..7297fd032b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/file_util.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/file_util.h @@ -31,7 +31,7 @@ #include #include #include "include/errorcode.h" -#include "utils/log_adapter.h" +#include "mindapi/base/logging.h" using mindspore::lite::RET_ERROR; using mindspore::lite::RET_OK; diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/float16.h b/mindspore/lite/tools/converter/adapter/dpico/common/float16.h new file mode 100644 index 0000000000..b2df561d5c --- /dev/null +++ b/mindspore/lite/tools/converter/adapter/dpico/common/float16.h @@ -0,0 +1,311 @@ +/** + * Copyright 2020-2022 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. + */ +#ifndef DPICO_COMMON_FLOAT16_H_ +#define DPICO_COMMON_FLOAT16_H_ + +#if defined(ENABLE_ARM32) || defined(ENABLE_ARM64) +// Built for lite and ARM +#include + +using float16 = float16_t; + +#else +#include +#include +#include +#include +#include +#include + +// Implement Float16 for mindspore, inspired by Eigen::half. +namespace mindspore { +class Float16 { + public: + static constexpr uint16_t value_mask = 0x7fff; + static constexpr uint16_t nan_value = 0x7e00; + static constexpr uint16_t inf_value = 0x7c00; + static constexpr uint16_t true_value = 0x3c00; + + union Union32 { + uint32_t u; + float f; + }; + + Float16() = default; + ~Float16() = default; + + Float16(const Float16 &other) noexcept = default; + Float16(Float16 &&other) noexcept = default; + + Float16 &operator=(const Float16 &other) noexcept = default; + Float16 &operator=(Float16 &&other) noexcept = default; + + static Float16 FromRaw(uint16_t v) { + Float16 f; + f.value_ = v; + return f; + } + + explicit Float16(float f) : value_(FromFloat32(f)) {} + explicit Float16(bool b) : value_(b ? true_value : 0) {} + template + explicit Float16(const T &v) : value_(FromFloat32(static_cast(v))) {} + + uint16_t int_value() const { return value_; } + + explicit operator bool() const { return (value_ & value_mask) != 0; } + explicit operator float() const { return ToFloat32(*this); } + explicit operator double() const { return static_cast(ToFloat32(*this)); } + explicit operator int8_t() const { return static_cast(ToFloat32(*this)); } + explicit operator uint8_t() const { return static_cast(ToFloat32(*this)); } + explicit operator int16_t() const { return static_cast(ToFloat32(*this)); } + explicit operator uint16_t() const { return static_cast(ToFloat32(*this)); } + explicit operator int32_t() const { return static_cast(ToFloat32(*this)); } + explicit operator uint32_t() const { return static_cast(ToFloat32(*this)); } + explicit operator int64_t() const { return static_cast(ToFloat32(*this)); } + explicit operator uint64_t() const { return static_cast(ToFloat32(*this)); } + + Float16 &operator+=(const Float16 &b) { + value_ = FromFloat32(ToFloat32(*this) + ToFloat32(b)); + return *this; + } + + Float16 &operator-=(const Float16 &b) { + value_ = FromFloat32(ToFloat32(*this) - ToFloat32(b)); + return *this; + } + + Float16 &operator*=(const Float16 &b) { + value_ = FromFloat32(ToFloat32(*this) * ToFloat32(b)); + return *this; + } + + Float16 &operator/=(const Float16 &b) { + value_ = FromFloat32(ToFloat32(*this) / ToFloat32(b)); + return *this; + } + + static float ToFloat32(const Float16 &f16) { + constexpr Union32 magic = {.u = 113 << 23}; + constexpr uint32_t exponent_adjust = ((127 - 15) << 23); + constexpr uint32_t inf_extra_exp_adjust = ((128 - 16) << 23); + constexpr uint32_t zero_extra_exp_adjust = (1 << 23); + constexpr uint32_t sign_mask = 0x8000; + constexpr unsigned int shifted_exp = (0x7c00 << 13); // Exponent mask after shift. + constexpr unsigned int exponent_bits = 13; + constexpr unsigned int sign_bit_shift = 16; + // Exponent/mantissa bits. + Union32 f32; + f32.u = (static_cast(f16.value_ & value_mask) << exponent_bits); + // Just the exponent. + unsigned int exp = (shifted_exp & f32.u); + f32.u += exponent_adjust; + // Handle exponent special cases. + if (exp == shifted_exp) { + // Inf/NaN, extra exp adjust. + f32.u += inf_extra_exp_adjust; + } else if (exp == 0) { + // Zero/Denormal, extra exp adjust and renormalize. + f32.u += zero_extra_exp_adjust; + f32.f -= magic.f; + } + // Set sign bit. + f32.u |= ((f16.value_ & sign_mask) << sign_bit_shift); + return f32.f; + } + + private: + static uint16_t FromFloat32(float f32) { + constexpr uint32_t magic = {113 << 23}; + constexpr Union32 f32infty = {.u = 255 << 23}; + constexpr Union32 f16max = {.u = (127 + 16) << 23}; + constexpr Union32 denorm_magic = {.u = ((127 - 15) + (23 - 10) + 1) << 23}; + constexpr unsigned int exponent_bits = 13; + constexpr unsigned int sign_bit_shift = 16; + constexpr unsigned int sign_mask = 0x80000000u; + constexpr uint32_t rouding_bias_part1 = ((unsigned int)(15 - 127) << 23) + 0xfff; + + Union32 f; + f.f = f32; + unsigned int sign = f.u & sign_mask; + f.u ^= sign; + uint16_t result = 0; + + // NOTE all the integer compares in this function can be safely + // compiled into signed compares since all operands are below + // 0x80000000. Important if you want fast straight SSE2 code + // (since there's no unsigned PCMPGTD). + if (f.u >= f16max.u) { + // Result is Inf or NaN (all exponent bits set). + result = (f.u > f32infty.u) ? nan_value : inf_value; + } else if (f.u < magic) { + // (De)normalized number or zero; resulting FP16 is subnormal or zero. + // Use a magic value to align our 10 mantissa bits at the bottom of + // the float. as long as FP addition is round-to-nearest-even this + // just works. + f.f += denorm_magic.f; + // And one integer subtract of the bias later, we have our final float! + result = static_cast(f.u - denorm_magic.u); + } else { + // Resulting mantissa is odd. + unsigned int mant_odd = (f.u >> exponent_bits) & 1; + // Update exponent, rounding bias part 1; + f.u += rouding_bias_part1; + // Rounding bias part 2; + f.u += mant_odd; + // Take the bits! + result = static_cast(f.u >> exponent_bits); + } + // Set sign bit. + result |= static_cast(sign >> sign_bit_shift); + return result; + } + + uint16_t value_; +}; + +inline Float16 operator+(const Float16 &a, const Float16 &b) { + return Float16(static_cast(a) + static_cast(b)); +} + +inline Float16 operator*(const Float16 &a, const Float16 &b) { + return Float16(static_cast(a) * static_cast(b)); +} + +inline Float16 operator-(const Float16 &a, const Float16 &b) { + return Float16(static_cast(a) - static_cast(b)); +} + +inline Float16 operator/(const Float16 &a, const Float16 &b) { + return Float16(static_cast(a) / static_cast(b)); +} + +// Division by an size_t. Do it in full float precision to avoid +// accuracy issues in converting the denominator to float16. +inline Float16 operator/(const Float16 &a, size_t b) { return Float16(static_cast(a) / static_cast(b)); } + +inline Float16 operator-(const Float16 &a) { + constexpr uint16_t sign_mask = 0x8000; + return Float16::FromRaw(a.int_value() ^ sign_mask); +} + +inline bool operator==(const Float16 &a, const Float16 &b) { + return std::equal_to()(static_cast(a), static_cast(b)); +} + +inline bool operator!=(const Float16 &a, const Float16 &b) { + return std::not_equal_to()(static_cast(a), static_cast(b)); +} + +inline bool operator<(const Float16 &a, const Float16 &b) { return static_cast(a) < static_cast(b); } +inline bool operator<=(const Float16 &a, const Float16 &b) { return static_cast(a) <= static_cast(b); } +inline bool operator>(const Float16 &a, const Float16 &b) { return static_cast(a) > static_cast(b); } +inline bool operator>=(const Float16 &a, const Float16 &b) { return static_cast(a) >= static_cast(b); } + +inline std::ostream &operator<<(std::ostream &os, const Float16 &v) { return (os << static_cast(v)); } + +} // namespace mindspore + +using float16 = mindspore::Float16; + +namespace std { +template <> +struct hash { + std::size_t operator()(const float16 &f16) const noexcept { return static_cast(f16.int_value()); } +}; + +template <> +struct numeric_limits { + static constexpr bool is_specialized = true; + static constexpr bool is_signed = true; + static constexpr bool is_integer = false; + static constexpr bool is_exact = false; + static constexpr bool has_infinity = true; + static constexpr bool has_quiet_NaN = true; + static constexpr bool has_signaling_NaN = true; + static constexpr std::float_denorm_style has_denorm = std::denorm_present; + static constexpr bool has_denorm_loss = false; + static constexpr std::float_round_style round_style = std::round_to_nearest; + static constexpr bool is_iec559 = false; + static constexpr bool is_bounded = false; + static constexpr bool is_modulo = false; + static constexpr int digits = 11; + static constexpr int digits10 = 3; + static constexpr int max_digits10 = 5; + static constexpr int radix = 2; + static constexpr int min_exponent = -13; + static constexpr int min_exponent10 = -4; + static constexpr int max_exponent = 16; + static constexpr int max_exponent10 = 4; + static constexpr bool traps = true; + static constexpr bool tinyness_before = false; + + static constexpr uint16_t raw_min = 0x400; + static constexpr uint16_t raw_max = 0x7bff; + static constexpr uint16_t raw_lowest = 0xfbff; + static constexpr uint16_t raw_epsilon = 0x0800; + static constexpr float round_error_value = 0.5; + + static float16(min)() noexcept { return float16::FromRaw(raw_min); } + static float16(max)() noexcept { return float16::FromRaw(raw_max); } + static float16 lowest() noexcept { return float16::FromRaw(raw_lowest); } + static float16 epsilon() noexcept { return float16::FromRaw(raw_epsilon); } + static float16 round_error() noexcept { return float16(round_error_value); } + static float16 infinity() noexcept { return float16::FromRaw(float16::inf_value); } + static float16 quiet_NaN() noexcept { return float16::FromRaw(float16::nan_value); } + static float16 signaling_NaN() noexcept { return float16::FromRaw(float16::nan_value); } + static float16 denorm_min() noexcept { return float16::FromRaw(1); } +}; + +// If std::numeric_limits is specialized, should also specialize +// std::numeric_limits, std::numeric_limits, and +// std::numeric_limits +// https://stackoverflow.com/a/16519653/ +template <> +struct numeric_limits : private numeric_limits {}; +template <> +struct numeric_limits : private numeric_limits {}; +template <> +struct numeric_limits : private numeric_limits {}; +} // namespace std + +// Implements standard math functions for float16. +inline bool(isinf)(const float16 &a) { return (a.int_value() & float16::value_mask) == float16::inf_value; } +inline bool(isnan)(const float16 &a) { return (a.int_value() & float16::value_mask) > float16::inf_value; } +inline bool(isfinite)(const float16 &a) { return !(isinf(a)) && !(isnan(a)); } +inline float16 abs(const float16 &a) { return float16::FromRaw(a.int_value() & float16::value_mask); } +inline float16 exp(const float16 &a) { return float16(::expf(static_cast(a))); } +inline float16 log(const float16 &a) { return float16(::logf(static_cast(a))); } +inline float16 log1p(const float16 &a) { return float16(::log1pf(static_cast(a))); } +inline float16 log10(const float16 &a) { return float16(::log10f(static_cast(a))); } +inline float16 sqrt(const float16 &a) { return float16(::sqrtf(static_cast(a))); } +inline float16 sin(const float16 &a) { return float16(::sinf(static_cast(a))); } +inline float16 cos(const float16 &a) { return float16(::cosf(static_cast(a))); } +inline float16 tan(const float16 &a) { return float16(::tanf(static_cast(a))); } +inline float16 tanh(const float16 &a) { return float16(::tanhf(static_cast(a))); } +inline float16 floor(const float16 &a) { return float16(::floorf(static_cast(a))); } +inline float16 ceil(const float16 &a) { return float16(::ceilf(static_cast(a))); } +inline float16(min)(const float16 &a, const float16 &b) { return b < a ? b : a; } +inline float16(max)(const float16 &a, const float16 &b) { return a < b ? b : a; } +inline float16 pow(const float16 &a, const float16 &b) { + return float16(::powf(static_cast(a), static_cast(b))); +} + +#endif // ENABLE_ARM32 || ENABLE_ARM64 + +inline float half_to_float(const float16 &h) { return static_cast(h); } + +#endif // DPICO_COMMON_FLOAT16_H_ diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.cc b/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.cc index 1226f756f6..ff67462b71 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.cc @@ -17,6 +17,10 @@ #include "common/format_utils.h" #include #include +#include "ops/tuple_get_item.h" +#include "ops/depend.h" +#include "ops/make_tuple.h" +#include "ops/return.h" #include "ops/batch_norm.h" #include "ops/batch_to_space.h" #include "ops/bias_add.h" @@ -54,11 +58,11 @@ const std::set kAssignedFormatOpSet = { const std::set &GetAssignedFormatOpSet() { return kAssignedFormatOpSet; } -bool IsSpecialType(const mindspore::CNodePtr &cnode) { - return CheckPrimitiveType(cnode, mindspore::prim::kPrimTupleGetItem) || - CheckPrimitiveType(cnode, mindspore::prim::kPrimDepend) || - CheckPrimitiveType(cnode, mindspore::prim::kPrimMakeTuple) || - CheckPrimitiveType(cnode, mindspore::prim::kPrimReturn); +bool IsSpecialType(const api::CNodePtr &cnode) { + return CheckPrimitiveType(cnode, api::MakeShared()) || + CheckPrimitiveType(cnode, api::MakeShared()) || + CheckPrimitiveType(cnode, api::MakeShared()) || + CheckPrimitiveType(cnode, api::MakeShared()); } std::string FormatEnumToString(mindspore::Format format) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.h b/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.h index aae4a0592b..f6f813f604 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/format_utils.h @@ -26,7 +26,7 @@ namespace mindspore { namespace dpico { const std::set &GetAssignedFormatOpSet(); -bool IsSpecialType(const mindspore::CNodePtr &cnode); +bool IsSpecialType(const api::CNodePtr &cnode); std::string FormatEnumToString(mindspore::Format format); } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/infer_util.cc b/mindspore/lite/tools/converter/adapter/dpico/common/infer_util.cc index 2b7b51ee43..efc498797e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/infer_util.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/infer_util.cc @@ -17,7 +17,7 @@ #include "common/infer_util.h" #include #include -#include "utils/log_adapter.h" +#include "mindapi/base/logging.h" #include "include/errorcode.h" using mindspore::lite::RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/op_attr.h b/mindspore/lite/tools/converter/adapter/dpico/common/op_attr.h index 7d29f7e49b..0ea6115d43 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/op_attr.h +++ b/mindspore/lite/tools/converter/adapter/dpico/common/op_attr.h @@ -32,7 +32,21 @@ constexpr auto kBlockWidth = "block_width"; constexpr auto kChannelShared = "channel_shared"; constexpr auto kCoeffs = "coeffs"; constexpr auto kCustomName = "custom_"; +constexpr auto kDetectionBackgroundLabelId = "detection_background_label_id"; +constexpr auto kDetectionBiasVec = "detection_bias_vec"; +constexpr auto kDetectionCalcMode = "detection_calc_mode"; +constexpr auto kDetectionClipBbox = "detection_clip_bbox"; +constexpr auto kDetectionCodeType = "detection_code_type"; +constexpr auto kDetectionMultiClassSorting = "detection_multi_class_sorting"; constexpr auto kDetectionOutputParam = "detection_output_param"; +constexpr auto kDetectionOutputParamSize = "detection_output_param_size"; +constexpr auto kDetectionProposalParamType = "detection_proposal_param_type"; +constexpr auto kDetectionReportFlag = "detection_report_flag"; +constexpr auto kDetectionShareLocation = "detection_share_location"; +constexpr auto kDetectionShareVariance = "detection_share_variance"; +constexpr auto kDetectionTop = "detection_top"; +constexpr auto kDetectionTopK = "detection_top_k"; +constexpr auto kDetectionVarianceVec = "detection_variance_vec"; constexpr auto kDecBBoxParam = "decbbox_param"; constexpr auto kDim1 = "dim_1"; constexpr auto kDim2 = "dim_2"; diff --git a/mindspore/lite/tools/converter/adapter/dpico/common/string_util.cc b/mindspore/lite/tools/converter/adapter/dpico/common/string_util.cc index b75f24e494..2e569d2bf6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/common/string_util.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/common/string_util.cc @@ -19,7 +19,7 @@ #include #include #include -#include "utils/log_adapter.h" +#include "mindapi/base/logging.h" namespace mindspore { namespace dpico { diff --git a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_extract_infer.cc b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_extract_infer.cc index fc6edf5956..ffae02feff 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_extract_infer.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_extract_infer.cc @@ -23,7 +23,7 @@ #include "utils/log_adapter.h" #include "common/infer_util.h" #include "include/errorcode.h" -#include "ops/op_utils.h" +#include "ops/op_name.h" #include "include/registry/register_kernel_interface.h" using mindspore::kernel::KernelInterface; diff --git a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_recurrent_infer.cc b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_recurrent_infer.cc index e3cdfa6738..dc15029660 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_recurrent_infer.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_recurrent_infer.cc @@ -31,7 +31,7 @@ namespace mindspore { namespace kernel { namespace { constexpr size_t kGateNum2 = 2; -} +} // namespace std::shared_ptr DpicoRecurrentInferCreater() { std::shared_ptr infer = std::make_shared(); if (infer == nullptr) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_roi_align_infer.cc b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_roi_align_infer.cc index d18557f941..cf2eaa7c3c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_roi_align_infer.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_roi_align_infer.cc @@ -25,6 +25,7 @@ #include "common/infer_util.h" #include "include/errorcode.h" #include "include/registry/register_kernel_interface.h" +#include "third_party/securec/include/securec.h" using mindspore::kernel::KernelInterface; using mindspore::lite::RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_spp_infer.cc b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_spp_infer.cc index c660f56085..a41928fd3a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_spp_infer.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_spp_infer.cc @@ -18,12 +18,12 @@ #include #include #include +#include #include "common/op_enum.h" #include "common/op_attr.h" #include "utils/log_adapter.h" #include "common/infer_util.h" #include "include/errorcode.h" -#include "ops/op_utils.h" #include "include/registry/register_kernel_interface.h" using mindspore::kernel::KernelInterface; @@ -95,7 +95,7 @@ Status DpicoSppInterface::Infer(std::vector *inputs, std::v MS_LOG(ERROR) << "input_shape should have 4 dims, but in fact it's " << input_shape.size(); return kLiteError; } - ShapeVector output_shape; + std::vector output_shape; output_shape.push_back(input_shape.at(0)); int64_t output_planes_size = 0; for (int64_t i = 0; i < pyramid_height; i++) { // spp output plane size is 1 * 1, 2 * 2, 4 * 4, ... diff --git a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_upsample_infer.cc b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_upsample_infer.cc index 56fcf4d2b2..0d93c94388 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_upsample_infer.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/infer/dpico_upsample_infer.cc @@ -24,7 +24,7 @@ #include "utils/log_adapter.h" #include "common/infer_util.h" #include "include/errorcode.h" -#include "ops/op_utils.h" +#include "ops/op_name.h" #include "include/registry/register_kernel_interface.h" using mindspore::kernel::KernelInterface; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.cc index 1ad169bec6..4eb8311b42 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS AbsMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS AbsMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.h index b322945323..9e7cad77aa 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/abs_mapper.h @@ -28,8 +28,8 @@ class AbsMapper : public OpMapper { public: AbsMapper() : OpMapper("Abs") {} ~AbsMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.cc index ffa72f2d08..494b385b50 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.cc @@ -20,7 +20,6 @@ #include #include #include "common/op_attr.h" -#include "ops/op_utils.h" #include "ops/fusion/activation.h" #include "op/relu_operator.h" #include "op/sigmoid_operator.h" @@ -33,7 +32,7 @@ namespace mindspore { namespace dpico { namespace { constexpr float kNum6 = 6.0; -std::unique_ptr ReluMapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr ReluMapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "relu operator is nullptr. "; @@ -46,7 +45,7 @@ std::unique_ptr ReluMapFunc(const std::shared_ptr Relu6MapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr Relu6MapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "relu operator is nullptr. "; @@ -61,7 +60,7 @@ std::unique_ptr Relu6MapFunc(const std::shared_ptr HardTanhMapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr HardTanhMapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "relu operator is nullptr. "; @@ -72,7 +71,7 @@ std::unique_ptr HardTanhMapFunc(const std::shared_ptr SigmoidMapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr SigmoidMapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "sigmoid operator is nullptr. "; @@ -82,7 +81,7 @@ std::unique_ptr SigmoidMapFunc(const std::shared_ptr TanhMapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr TanhMapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "tanh operator is nullptr. "; @@ -91,7 +90,7 @@ std::unique_ptr TanhMapFunc(const std::shared_ptr HswishMapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr HswishMapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "hswish operator is nullptr. "; @@ -100,12 +99,12 @@ std::unique_ptr HswishMapFunc(const std::shared_ptrSetOpType(mapper::OpType::HSWISH); if (activation_prim->GetAttr(dpico::kNegativeSlope) != nullptr) { static_cast(activation_operator.get()) - ->SetHswishSlope(GetValue(activation_prim->GetAttr(dpico::kNegativeSlope))); + ->SetHswishSlope(api::GetValue(activation_prim->GetAttr(dpico::kNegativeSlope))); } return std::move(activation_operator); } -std::unique_ptr EluMapFunc(const std::shared_ptr &activation_prim) { +std::unique_ptr EluMapFunc(const api::SharedPtr &activation_prim) { auto activation_operator = std::make_unique(); if (activation_operator == nullptr) { MS_LOG(ERROR) << "elu operator is nullptr. "; @@ -114,13 +113,13 @@ std::unique_ptr EluMapFunc(const std::shared_ptrSetOpType(mapper::OpType::ELU); if (activation_prim->GetAttr(ops::kAlpha) != nullptr) { static_cast(activation_operator.get()) - ->SetEluAlpha(GetValue(activation_prim->GetAttr(ops::kAlpha))); + ->SetEluAlpha(api::GetValue(activation_prim->GetAttr(ops::kAlpha))); } return std::move(activation_operator); } using ActivationMapFunc = - std::unique_ptr (*)(const std::shared_ptr &activation_prim); + std::unique_ptr (*)(const api::SharedPtr &activation_prim); const std::unordered_map kActivationMapFuncs = { {ActivationType::RELU, &ReluMapFunc}, {ActivationType::LEAKY_RELU, &ReluMapFunc}, {ActivationType::RELU6, &Relu6MapFunc}, {ActivationType::HARD_TANH, &HardTanhMapFunc}, @@ -128,13 +127,13 @@ const std::unordered_map kActivationMapFuncs {ActivationType::HSWISH, &HswishMapFunc}, {ActivationType::ELU, &EluMapFunc}}; } // namespace -STATUS ActivationMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS ActivationMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto activation_prim = utils::cast>(prim); + auto activation_prim = api::utils::cast>(prim); MS_ASSERT(activation_prim != nullptr); auto activation_type = activation_prim->get_activation_type(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.h index d7ab370120..1230192bb6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/activation_mapper.h @@ -28,8 +28,8 @@ class ActivationMapper : public OpMapper { public: ActivationMapper() : OpMapper("Activation") {} ~ActivationMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.cc index d4a7d83a8c..403bf6e25f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.cc @@ -18,20 +18,19 @@ #include #include #include -#include "ops/op_utils.h" #include "common/anf_util.h" #include "ops/fusion/arg_max_fusion.h" #include "op/argmax_operator.h" namespace mindspore { namespace dpico { -STATUS ArgMaxMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ArgMaxMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto argmax_prim = utils::cast>(prim); + auto argmax_prim = api::utils::cast>(prim); MS_ASSERT(argmax_prim != nullptr); auto argmax_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.h index c665c4f0a2..d6734ca3dd 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/argmax_mapper.h @@ -28,8 +28,8 @@ class ArgMaxMapper : public OpMapper { public: ArgMaxMapper() : OpMapper("ArgMax") {} ~ArgMaxMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.cc index 3fc517535b..79fb67a231 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.cc @@ -20,6 +20,7 @@ #include #include #include +#include "ops/neg.h" #include "common/anf_util.h" #include "op/binary_math_operator.h" @@ -40,8 +41,8 @@ const std::unordered_map kArithmeticOpMap = { {"X_DIV_Y", mapper::BinaryMathOp::X_DIV_Y_OP}, {"X_LOG_Y", mapper::BinaryMathOp::X_LOG_Y_OP}}; } // namespace -STATUS ArithmeticMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS ArithmeticMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -71,7 +72,7 @@ STATUS ArithmeticMapper::Map(const CNodePtr &cnode, std::vector auto binary_math_op = kArithmeticOpMap.at(op_type_name); arithmetic_operator->SetBinaryMathOp(binary_math_op); - if (CheckPrimitiveType(cnode, prim::kPrimNeg)) { + if (CheckPrimitiveType(cnode, api::MakeShared())) { arithmetic_operator->PushOfflineArgs(std::make_pair(std::vector{-1.0}, std::vector{})); arithmetic_operator->PushOfflineArgs(std::make_pair(std::vector{}, std::vector{})); } else { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.h index 8dc0bc62c5..fa552b9245 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/arithmetic_mapper.h @@ -28,8 +28,8 @@ class ArithmeticMapper : public OpMapper { public: ArithmeticMapper() : OpMapper("Arithmetic") {} ~ArithmeticMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.cc index f9425a5dfd..84a9b334fc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.cc @@ -19,7 +19,8 @@ #include #include #include -#include "ops/op_utils.h" +#include +#include "ops/batch_norm.h" #include "common/anf_util.h" #include "common/op_enum.h" #include "op/batch_norm_operator.h" @@ -28,7 +29,7 @@ namespace mindspore { namespace dpico { namespace { // BatchNorm: {BNMeanIndex:2, BNVarIndex:3, ScaleFactorIndex:4} -STATUS SetBnDataInfo(const CNodePtr &cnode, mapper::BatchNormOperator *batch_norm_operator) { +STATUS SetBnDataInfo(const api::CNodePtr &cnode, mapper::BatchNormOperator *batch_norm_operator) { if (batch_norm_operator == nullptr) { MS_LOG(ERROR) << "batch_norm_operator is nullptr."; return RET_ERROR; @@ -36,13 +37,13 @@ STATUS SetBnDataInfo(const CNodePtr &cnode, mapper::BatchNormOperator *batch_nor for (size_t i = 2; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); MS_ASSERT(input_node != nullptr); - auto param_node = input_node->cast(); + auto param_node = input_node->cast(); if (param_node == nullptr || !param_node->has_default()) { continue; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); if (i == kInputIndex2) { batch_norm_operator->SetBnMeanDataPtr(data); @@ -67,7 +68,7 @@ STATUS SetBnDataInfo(const CNodePtr &cnode, mapper::BatchNormOperator *batch_nor return RET_OK; } // FusedBatchNorm: {ScaleIndex:2, BiasIndex:3, MeanIndex:4, VarIndex:5} -STATUS SetFusedBnDataInfo(const CNodePtr &cnode, mapper::BatchNormOperator *batch_norm_operator) { +STATUS SetFusedBnDataInfo(const api::CNodePtr &cnode, mapper::BatchNormOperator *batch_norm_operator) { if (batch_norm_operator == nullptr) { MS_LOG(ERROR) << "batch_norm_operator is nullptr."; return RET_ERROR; @@ -75,13 +76,13 @@ STATUS SetFusedBnDataInfo(const CNodePtr &cnode, mapper::BatchNormOperator *batc for (size_t i = 2; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); MS_ASSERT(input_node != nullptr); - auto param_node = input_node->cast(); + auto param_node = input_node->cast(); if (param_node == nullptr || !param_node->has_default()) { continue; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); if (i == kInputIndex2) { batch_norm_operator->SetBnScalePtr(data); @@ -109,8 +110,8 @@ STATUS SetFusedBnDataInfo(const CNodePtr &cnode, mapper::BatchNormOperator *batc return RET_OK; } } // namespace -STATUS BatchNormMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS BatchNormMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -128,19 +129,19 @@ STATUS BatchNormMapper::Map(const CNodePtr &cnode, std::vector batch_norm_operator->SetOpType(mapper::OpType::BN); if (prim->GetAttr(ops::kEpsilon) != nullptr) { - batch_norm_operator->SetBnEps(GetValue(prim->GetAttr(ops::kEpsilon))); + batch_norm_operator->SetBnEps(api::GetValue(prim->GetAttr(ops::kEpsilon))); } if (prim->GetAttr(ops::kMomentum) != nullptr) { - auto momentum = GetValue(prim->GetAttr(ops::kMomentum)); + auto momentum = api::GetValue(prim->GetAttr(ops::kMomentum)); const float default_momentum_value = 0.9; - if (fabs(momentum - default_momentum_value) > std::numeric_limits::epsilon()) { + if (std::fabs(momentum - default_momentum_value) > std::numeric_limits::epsilon()) { MS_LOG(INFO) << cnode->fullname_with_scope() << "'s momentum attr value " << momentum << " is not equal to mapper default value 0.9. Note that mapper will ignore this value."; } } - if (CheckPrimitiveType(cnode, prim::kPrimBatchNorm)) { + if (CheckPrimitiveType(cnode, api::MakeShared())) { if (SetBnDataInfo(cnode, batch_norm_operator.get()) != RET_OK) { MS_LOG(ERROR) << "set bn data info failed."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.h index 309c14bfbc..d887d28bbf 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/batch_norm_mapper.h @@ -28,8 +28,8 @@ class BatchNormMapper : public OpMapper { public: BatchNormMapper() : OpMapper("BatchNorm") {} ~BatchNormMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.cc index 5de03df1f7..724d787003 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS BiLstmMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS BiLstmMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -42,13 +42,13 @@ STATUS BiLstmMapper::Map(const CNodePtr &cnode, std::vector *ba } if (prim->GetAttr(kNumOutput) != nullptr) { - bi_lstm_operator->SetRecurrentNumOutput(GetValue(prim->GetAttr(kNumOutput))); + bi_lstm_operator->SetRecurrentNumOutput(static_cast(api::GetValue(prim->GetAttr(kNumOutput)))); } if (prim->GetAttr(kExposeHidden) != nullptr) { - bi_lstm_operator->SetRecurrentExposeHidden(GetValue(prim->GetAttr(kExposeHidden))); + bi_lstm_operator->SetRecurrentExposeHidden(api::GetValue(prim->GetAttr(kExposeHidden))); } if (prim->GetAttr(kOutputChannel) != nullptr) { - bi_lstm_operator->SetOutputChannel(GetValue(prim->GetAttr(kOutputChannel))); + bi_lstm_operator->SetOutputChannel(static_cast(api::GetValue(prim->GetAttr(kOutputChannel)))); } bi_lstm_operator->SetRecurrentContFlag(true); if (SetRecurrentDataInfo(cnode, bi_lstm_operator.get()) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.h index d90232a812..a0739f6ace 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/bi_lstm_mapper.h @@ -28,8 +28,8 @@ class BiLstmMapper : public OpMapper { public: BiLstmMapper() : OpMapper("BiLstm") {} ~BiLstmMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.cc index bc3bc865db..5a49fb6574 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.cc @@ -19,7 +19,7 @@ #include #include #include -#include "ops/op_utils.h" +#include "ops/bias_add.h" #include "common/anf_util.h" #include "common/op_attr.h" #include "common/op_enum.h" @@ -28,7 +28,7 @@ namespace mindspore { namespace dpico { namespace { -STATUS SetBiasDataInfo(const CNodePtr &cnode, mapper::BiasOperator *bias_operator) { +STATUS SetBiasDataInfo(const api::CNodePtr &cnode, mapper::BiasOperator *bias_operator) { if (bias_operator == nullptr) { MS_LOG(ERROR) << "bias_operator is nullptr."; return RET_ERROR; @@ -36,14 +36,14 @@ STATUS SetBiasDataInfo(const CNodePtr &cnode, mapper::BiasOperator *bias_operato if (cnode->inputs().size() == kInputIndex3) { auto input_anode = cnode->input(kInputIndex2); MS_ASSERT(input_anode != nullptr); - auto param_node = input_anode->cast(); + auto param_node = input_anode->cast(); if (param_node == nullptr || !param_node->has_default()) { MS_LOG(DEBUG) << "only parameter node needs to set BiasPtr"; return RET_OK; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); bias_operator->SetBiasPtr(data); ShapeVector shape_vector; @@ -64,8 +64,8 @@ STATUS SetBiasDataInfo(const CNodePtr &cnode, mapper::BiasOperator *bias_operato return RET_OK; } } // namespace -STATUS BiasMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS BiasMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -83,9 +83,9 @@ STATUS BiasMapper::Map(const CNodePtr &cnode, std::vector *base } if (prim->GetAttr(ops::kAxis) != nullptr) { - bias_operator->SetAxis(GetValue(prim->GetAttr(ops::kAxis))); - } else if (CheckPrimitiveType(cnode, prim::kPrimBiasAdd)) { - auto format = GetValue(prim->GetAttr(ops::kFormat)); + bias_operator->SetAxis(static_cast(api::GetValue(prim->GetAttr(ops::kAxis)))); + } else if (CheckPrimitiveType(cnode, api::MakeShared())) { + auto format = api::GetValue(prim->GetAttr(ops::kFormat)); if (format == mindspore::NCHW) { bias_operator->SetAxis(1); } else if (format == mindspore::NHWC) { @@ -96,7 +96,7 @@ STATUS BiasMapper::Map(const CNodePtr &cnode, std::vector *base } } if (prim->GetAttr(dpico::kNumAxes) != nullptr) { - bias_operator->SetBiasNumAxes(GetValue(prim->GetAttr(dpico::kNumAxes))); + bias_operator->SetBiasNumAxes(static_cast(api::GetValue(prim->GetAttr(dpico::kNumAxes)))); } if (SetBiasDataInfo(cnode, bias_operator.get()) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.h index 448be98e76..c6c9487e70 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/bias_mapper.h @@ -28,8 +28,8 @@ class BiasMapper : public OpMapper { public: BiasMapper() : OpMapper("Bias") {} ~BiasMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.cc index ddeb777808..cef1761b59 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.cc @@ -23,8 +23,8 @@ namespace mindspore { namespace dpico { -STATUS BnllMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS BnllMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.h index 50c4b49086..8a9674e41a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/bnll_mapper.h @@ -28,8 +28,8 @@ class BnllMapper : public OpMapper { public: BnllMapper() : OpMapper("Bnll") {} ~BnllMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.cc index a233e19687..6dfec65da9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.cc @@ -18,14 +18,13 @@ #include #include #include -#include "ops/op_utils.h" #include "ops/cast.h" #include "op/cast_operator.h" namespace mindspore { namespace dpico { -STATUS CastMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS CastMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.h index d4d3c3a578..23aff078dd 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/cast_mapper.h @@ -28,8 +28,8 @@ class CastMapper : public OpMapper { public: CastMapper() : OpMapper("Cast") {} ~CastMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.cc index f04f39fcd4..a7874ab752 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.cc @@ -18,19 +18,18 @@ #include #include #include -#include "ops/op_utils.h" #include "ops/clip.h" #include "op/clip_operator.h" namespace mindspore { namespace dpico { -STATUS ClipMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ClipMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto clip_prim = utils::cast>(prim); + auto clip_prim = api::utils::cast>(prim); MS_ASSERT(clip_prim != nullptr); auto clip_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.h index 8cb9beed02..09adb46d14 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/clip_mapper.h @@ -28,8 +28,8 @@ class ClipMapper : public OpMapper { public: ClipMapper() : OpMapper("Clip") {} ~ClipMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.cc index bb230bddaa..dfd6b2d8f9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.cc @@ -18,19 +18,18 @@ #include #include #include -#include "ops/op_utils.h" #include "ops/concat.h" #include "op/concat_operator.h" namespace mindspore { namespace dpico { -STATUS ConcatMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ConcatMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto concat_prim = utils::cast>(prim); + auto concat_prim = api::utils::cast>(prim); MS_ASSERT(concat_prim != nullptr); auto concat_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.h index a2e074d869..b88f8c3cf0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/concat_mapper.h @@ -28,8 +28,8 @@ class ConcatMapper : public OpMapper { public: ConcatMapper() : OpMapper("Concat") {} ~ConcatMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.cc index 4717dcff03..4949257d29 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.cc @@ -38,7 +38,7 @@ struct ConvAttr { int64_t out_channel{1}; PadMode pad_mode{PadMode::PAD}; }; -STATUS GetConvAttrFromPrimitive(ConvAttr *conv_attr, const std::shared_ptr &conv_prim) { +STATUS GetConvAttrFromPrimitive(ConvAttr *conv_attr, const api::SharedPtr &conv_prim) { if (conv_attr == nullptr) { MS_LOG(ERROR) << "conv_attr is nullptr."; return RET_ERROR; @@ -82,13 +82,13 @@ STATUS GetConvAttrFromPrimitive(ConvAttr *conv_attr, const std::shared_ptr *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ConvMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto conv_prim = utils::cast>(prim); + auto conv_prim = api::utils::cast>(prim); MS_ASSERT(conv_prim != nullptr); ConvAttr conv_attr; if (GetConvAttrFromPrimitive(&conv_attr, conv_prim) != RET_OK) { @@ -110,7 +110,7 @@ STATUS ConvMapper::Map(const CNodePtr &cnode, std::vector *base } std::unique_ptr conv_operator; - if (conv_prim->GetAttr(ops::kIsDepthWise) != nullptr && GetValue(conv_prim->GetAttr(ops::kIsDepthWise))) { + if (conv_prim->GetAttr(ops::kIsDepthWise) != nullptr && api::GetValue(conv_prim->GetAttr(ops::kIsDepthWise))) { conv_operator = std::make_unique(); if (conv_operator == nullptr) { MS_LOG(ERROR) << "conv_operator is nullptr."; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.h index dbe48e4c8c..b617810643 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/conv_mapper.h @@ -28,8 +28,8 @@ class ConvMapper : public OpMapper { public: ConvMapper() : OpMapper("Convolution") {} ~ConvMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.cc index eb0512723a..71decd105d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.cc @@ -19,19 +19,18 @@ #include #include #include -#include "ops/op_utils.h" #include "ops/crop.h" #include "op/crop_operator.h" namespace mindspore { namespace dpico { -STATUS CropMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS CropMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto crop_prim = utils::cast>(prim); + auto crop_prim = api::utils::cast>(prim); MS_ASSERT(crop_prim != nullptr); auto crop_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.h index d8172d7895..770713564c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/crop_mapper.h @@ -28,8 +28,8 @@ class CropMapper : public OpMapper { public: CropMapper() : OpMapper("Crop") {} ~CropMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.cc index dfea2d1af8..354389769b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.cc @@ -18,19 +18,22 @@ #include #include #include -#include "parser/detection_output_param_holder.h" #include "common/anf_util.h" #include "common/op_attr.h" #include "op/dec_bbox_operator.h" +#include "ops/custom.h" +#include "parser/detection_output_param_helper.h" namespace mindspore { namespace dpico { -STATUS DecBBoxMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS DecBBoxMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } + auto custom_prim = api::utils::cast>(prim); + MS_CHECK_TRUE_MSG(custom_prim != nullptr, RET_ERROR, "custom_prim is nullptr"); auto decbbox_operator = std::make_unique(); if (decbbox_operator == nullptr) { MS_LOG(ERROR) << "decbbox_operator is nullptr."; @@ -44,30 +47,31 @@ STATUS DecBBoxMapper::Map(const CNodePtr &cnode, std::vector *b decbbox_operator->SetOpType(mapper::OpType::DECBBOX); if (prim->GetAttr(kNumAnchors) != nullptr) { - decbbox_operator->SetNumAnchors(GetValue(prim->GetAttr(kNumAnchors))); + decbbox_operator->SetNumAnchors(static_cast(api::GetValue(prim->GetAttr(kNumAnchors)))); } if (prim->GetAttr(kNumBboxesPerGrid) != nullptr) { - decbbox_operator->SetNumBboxesPerGrid(GetValue(prim->GetAttr(kNumBboxesPerGrid))); + decbbox_operator->SetNumBboxesPerGrid( + static_cast(api::GetValue(prim->GetAttr(kNumBboxesPerGrid)))); } if (prim->GetAttr(kNumCoords) != nullptr) { - decbbox_operator->SetNumCoords(GetValue(prim->GetAttr(kNumCoords))); + decbbox_operator->SetNumCoords(static_cast(api::GetValue(prim->GetAttr(kNumCoords)))); } if (prim->GetAttr(kNumClasses) != nullptr) { - decbbox_operator->SetNumClasses(GetValue(prim->GetAttr(kNumClasses))); + decbbox_operator->SetNumClasses(static_cast(api::GetValue(prim->GetAttr(kNumClasses)))); } if (prim->GetAttr(kNumGridsHeight) != nullptr) { - decbbox_operator->SetNumGridsHeight(GetValue(prim->GetAttr(kNumGridsHeight))); + decbbox_operator->SetNumGridsHeight(static_cast(api::GetValue(prim->GetAttr(kNumGridsHeight)))); } if (prim->GetAttr(kNumGridsWidth) != nullptr) { - decbbox_operator->SetNumGridsWidth(GetValue(prim->GetAttr(kNumGridsWidth))); + decbbox_operator->SetNumGridsWidth(static_cast(api::GetValue(prim->GetAttr(kNumGridsWidth)))); } - if (prim->GetAttr(kDecBBoxParam) != nullptr) { - auto param_ptr = GetValue(prim->GetAttr(kDecBBoxParam)); - if (param_ptr == nullptr) { - MS_LOG(ERROR) << "decbbox param holder ptr is nullptr."; - return RET_ERROR; - } - decbbox_operator->SetDecBboxParam(param_ptr->GetDetectionOutputParam()); + std::vector param_vec; + if (GetDetectionOutputParamFromAttrs(¶m_vec, custom_prim) != RET_OK) { + MS_LOG(ERROR) << "get detection output param from attrs failed."; + return RET_ERROR; + } + if (param_vec.size() == 1) { + decbbox_operator->SetDecBboxParam(param_vec.at(0)); } base_operators->push_back(std::move(decbbox_operator)); return RET_OK; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.h index 1af6987813..11b84b0e35 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/decbbox_mapper.h @@ -29,8 +29,8 @@ class DecBBoxMapper : public OpMapper { public: DecBBoxMapper() : OpMapper("DecBBox") {} ~DecBBoxMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.cc index b6c6208681..bdd1814c2e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.cc @@ -28,13 +28,13 @@ constexpr int kDeconvKernelSize = 2; constexpr int kDeconvStrideSize = 2; constexpr int kDeconvDilationSize = 2; constexpr int kDeconvPadListSize = 4; -STATUS DeconvMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS DeconvMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto deconv_prim = utils::cast>(prim); + auto deconv_prim = api::utils::cast>(prim); MS_ASSERT(deconv_prim != nullptr); auto kernel_size = deconv_prim->get_kernel_size(); auto stride = deconv_prim->get_stride(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.h index a81603a4b4..3b7b5c7dc1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/deconv_mapper.h @@ -28,8 +28,8 @@ class DeconvMapper : public OpMapper { public: DeconvMapper() : OpMapper("Deconvolution") {} ~DeconvMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.cc index 33ae4b88ea..cbb1d1c306 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.cc @@ -18,19 +18,21 @@ #include #include #include -#include "parser/detection_output_param_holder.h" #include "common/anf_util.h" #include "common/op_attr.h" #include "op/detection_output_operator.h" +#include "parser/detection_output_param_helper.h" namespace mindspore { namespace dpico { -STATUS DetectionOutputMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS DetectionOutputMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } + auto custom_prim = api::utils::cast>(prim); + MS_CHECK_TRUE_MSG(custom_prim != nullptr, RET_ERROR, "custom_prim is nullptr"); auto detection_output_operator = std::make_unique(); if (detection_output_operator == nullptr) { MS_LOG(ERROR) << "detection_output_operator is nullptr."; @@ -44,32 +46,33 @@ STATUS DetectionOutputMapper::Map(const CNodePtr &cnode, std::vectorSetOpType(mapper::OpType::DETECTION_OUTPUT); if (prim->GetAttr(kNumAnchors) != nullptr) { - detection_output_operator->SetNumAnchors(GetValue(prim->GetAttr(kNumAnchors))); + detection_output_operator->SetNumAnchors(static_cast(api::GetValue(prim->GetAttr(kNumAnchors)))); } if (prim->GetAttr(kNumBboxesPerGrid) != nullptr) { - detection_output_operator->SetNumBboxesPerGrid(GetValue(prim->GetAttr(kNumBboxesPerGrid))); + detection_output_operator->SetNumBboxesPerGrid( + static_cast(api::GetValue(prim->GetAttr(kNumBboxesPerGrid)))); } if (prim->GetAttr(kNumCoords) != nullptr) { - detection_output_operator->SetNumCoords(GetValue(prim->GetAttr(kNumCoords))); + detection_output_operator->SetNumCoords(static_cast(api::GetValue(prim->GetAttr(kNumCoords)))); } if (prim->GetAttr(kNumClasses) != nullptr) { - detection_output_operator->SetNumClasses(GetValue(prim->GetAttr(kNumClasses))); + detection_output_operator->SetNumClasses(static_cast(api::GetValue(prim->GetAttr(kNumClasses)))); } if (prim->GetAttr(kNumGridsHeight) != nullptr) { - detection_output_operator->SetNumGridsHeight(GetValue(prim->GetAttr(kNumGridsHeight))); + detection_output_operator->SetNumGridsHeight( + static_cast(api::GetValue(prim->GetAttr(kNumGridsHeight)))); } if (prim->GetAttr(kNumGridsWidth) != nullptr) { - detection_output_operator->SetNumGridsWidth(GetValue(prim->GetAttr(kNumGridsWidth))); + detection_output_operator->SetNumGridsWidth( + static_cast(api::GetValue(prim->GetAttr(kNumGridsWidth)))); } - if (prim->GetAttr(kDetectionOutputParam) != nullptr) { - auto param_ptr_list = - GetValue>(prim->GetAttr(kDetectionOutputParam)); - std::vector param_vec{}; - (void)std::transform( - param_ptr_list.begin(), param_ptr_list.end(), std::back_inserter(param_vec), - [](const lite::DetectionOutputParamHolderPtr ¶m_ptr) { return param_ptr->GetDetectionOutputParam(); }); - detection_output_operator->SetDetectionOutputParamVec(param_vec); + + std::vector param_vec; + if (GetDetectionOutputParamFromAttrs(¶m_vec, custom_prim) != RET_OK) { + MS_LOG(ERROR) << "get detection output param from attrs failed."; + return RET_ERROR; } + detection_output_operator->SetDetectionOutputParamVec(param_vec); base_operators->push_back(std::move(detection_output_operator)); return RET_OK; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.h index f873f11b08..c0630eb718 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/detection_output_mapper.h @@ -28,8 +28,8 @@ class DetectionOutputMapper : public OpMapper { public: DetectionOutputMapper() : OpMapper("DetectionOutput") {} ~DetectionOutputMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.cc index 3c7b3db535..60bbb4b174 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.cc @@ -18,15 +18,15 @@ #include #include #include -#include "ops/op_utils.h" #include "common/op_attr.h" #include "common/anf_util.h" #include "op/eltwise_operator.h" +#include "mindapi/base/types.h" namespace mindspore { namespace dpico { -STATUS EltwiseMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS EltwiseMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -48,7 +48,7 @@ STATUS EltwiseMapper::Map(const CNodePtr &cnode, std::vector *b MS_LOG(ERROR) << "eltwise op should have mode attr. " << cnode->fullname_with_scope(); return RET_ERROR; } - auto eltwise_mode = static_cast(GetValue(prim->GetAttr(ops::kMode))); + auto eltwise_mode = static_cast(api::GetValue(prim->GetAttr(ops::kMode))); switch (eltwise_mode) { case mindspore::EltwiseMode::PROD: eltwise_operator->SetEltwiseOp(mapper::BinaryMathOp::MUL_OP); @@ -65,7 +65,7 @@ STATUS EltwiseMapper::Map(const CNodePtr &cnode, std::vector *b } if (prim->GetAttr(dpico::kCoeffs) != nullptr) { - eltwise_operator->SetEltCoeffVec(GetValue>(prim->GetAttr(dpico::kCoeffs))); + eltwise_operator->SetEltCoeffVec(api::GetValue>(prim->GetAttr(dpico::kCoeffs))); } base_operators->push_back(std::move(eltwise_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.h index 7d3abeea38..c8f4a0be44 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/eltwise_mapper.h @@ -28,8 +28,8 @@ class EltwiseMapper : public OpMapper { public: EltwiseMapper() : OpMapper("Eltwise") {} ~EltwiseMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.cc index beee63c57f..c4f8bfd264 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.cc @@ -20,19 +20,18 @@ #include #include "common/anf_util.h" #include "common/op_enum.h" -#include "ops/op_utils.h" #include "ops/fusion/exp_fusion.h" #include "op/exp_operator.h" namespace mindspore { namespace dpico { -STATUS ExpMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ExpMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto exp_prim = utils::cast>(prim); + auto exp_prim = api::utils::cast>(prim); MS_ASSERT(exp_prim != nullptr); auto exp_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.h index 8f85371e37..6248a1ef36 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/exp_mapper.h @@ -28,8 +28,8 @@ class ExpMapper : public OpMapper { public: ExpMapper() : OpMapper("Exp") {} ~ExpMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.cc index c8760d9b3f..56af99b039 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.cc @@ -18,15 +18,14 @@ #include #include #include -#include "ops/op_utils.h" #include "common/anf_util.h" #include "common/op_attr.h" #include "op/extract_operator.h" namespace mindspore { namespace dpico { -STATUS ExtractMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ExtractMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -44,13 +43,13 @@ STATUS ExtractMapper::Map(const CNodePtr &cnode, std::vector *b extract_operator->SetOpType(mapper::OpType::EXTRACT); if (prim->GetAttr(dpico::kSlicePointBegin) != nullptr) { - extract_operator->SetSlicePointBegin(GetValue(prim->GetAttr(dpico::kSlicePointBegin))); + extract_operator->SetSlicePointBegin(api::GetValue(prim->GetAttr(dpico::kSlicePointBegin))); } if (prim->GetAttr(dpico::kSlicePointEnd) != nullptr) { - extract_operator->SetSlicePointEnd(GetValue(prim->GetAttr(dpico::kSlicePointEnd))); + extract_operator->SetSlicePointEnd(api::GetValue(prim->GetAttr(dpico::kSlicePointEnd))); } if (prim->GetAttr(ops::kAxis) != nullptr) { - extract_operator->SetAxis(GetValue(prim->GetAttr(ops::kAxis))); + extract_operator->SetAxis(api::GetValue(prim->GetAttr(ops::kAxis))); } base_operators->push_back(std::move(extract_operator)); return RET_OK; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.h index 3ff801d9a2..909fa24202 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/extract_mapper.h @@ -28,8 +28,8 @@ class ExtractMapper : public OpMapper { public: ExtractMapper() : OpMapper("Extract") {} ~ExtractMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.cc index 5233c28956..0b9177dad8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.cc @@ -20,13 +20,12 @@ #include #include "common/op_attr.h" #include "common/anf_util.h" -#include "ops/op_utils.h" #include "op/fc_operator.h" namespace mindspore { namespace dpico { namespace { -STATUS SetNumOutput(const CNodePtr &cnode, const PrimitivePtr &prim, mapper::FcOperator *fc_operator) { +STATUS SetNumOutput(const api::CNodePtr &cnode, const api::PrimitivePtr &prim, mapper::FcOperator *fc_operator) { if (fc_operator == nullptr) { MS_LOG(ERROR) << "fc_operator is nullptr."; return RET_ERROR; @@ -42,7 +41,7 @@ STATUS SetNumOutput(const CNodePtr &cnode, const PrimitivePtr &prim, mapper::FcO } auto output_shape = output_shapes.at(0); if (prim->GetAttr(kNumOutput) != nullptr) { - auto num_output = GetValue(prim->GetAttr(kNumOutput)); + uint32_t num_output = api::GetValue(prim->GetAttr(kNumOutput)); if (output_shape.back() != num_output) { MS_LOG(ERROR) << "num output attr isn't matched with fc output shape."; return RET_ERROR; @@ -54,8 +53,8 @@ STATUS SetNumOutput(const CNodePtr &cnode, const PrimitivePtr &prim, mapper::FcO return RET_OK; } } // namespace -STATUS FCMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS FCMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -74,11 +73,11 @@ STATUS FCMapper::Map(const CNodePtr &cnode, std::vector *base_o fc_operator->SetOpType(mapper::OpType::INNERPRODUCT); if (prim->GetAttr(ops::kAxis) != nullptr) { - fc_operator->SetAxis(static_cast(GetValue(prim->GetAttr(ops::kAxis)))); + fc_operator->SetAxis(static_cast(api::GetValue(prim->GetAttr(ops::kAxis)))); } if (prim->GetAttr(ops::kTransposeB) != nullptr) { // note that this value of fc operator is opposite to kTransposeB - fc_operator->SetFcTransposeFlag(!GetValue(prim->GetAttr(ops::kTransposeB))); + fc_operator->SetFcTransposeFlag(!api::GetValue(prim->GetAttr(ops::kTransposeB))); } if (SetNumOutput(cnode, prim, fc_operator.get()) != RET_OK) { MS_LOG(ERROR) << "set num output failed."; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.h index 1dc2ccd43b..70ca54b278 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/fc_mapper.h @@ -28,8 +28,8 @@ class FCMapper : public OpMapper { public: FCMapper() : OpMapper("FC") {} ~FCMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.cc index 66d0c79d16..48ae0b7dcb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.cc @@ -25,13 +25,13 @@ namespace mindspore { namespace dpico { -STATUS FlattenMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS FlattenMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto flatten_prim = utils::cast>(prim); + auto flatten_prim = api::utils::cast>(prim); MS_ASSERT(flatten_prim != nullptr); auto flatten_operator = std::make_unique(); @@ -47,10 +47,11 @@ STATUS FlattenMapper::Map(const CNodePtr &cnode, std::vector *b flatten_operator->SetOpType(mapper::OpType::FLATTEN); if (flatten_prim->GetAttr(kStartAxis) != nullptr) { - flatten_operator->SetFlattenStartAxis(GetValue(flatten_prim->GetAttr(kStartAxis))); + flatten_operator->SetFlattenStartAxis( + static_cast(api::GetValue(flatten_prim->GetAttr(kStartAxis)))); } if (flatten_prim->GetAttr(kEndAxis) != nullptr) { - flatten_operator->SetFlattenEndAxis(GetValue(flatten_prim->GetAttr(kEndAxis))); + flatten_operator->SetFlattenEndAxis(static_cast(api::GetValue(flatten_prim->GetAttr(kEndAxis)))); } base_operators->push_back(std::move(flatten_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.h index eada0191ef..d8f4f30f60 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/flatten_mapper.h @@ -28,8 +28,8 @@ class FlattenMapper : public OpMapper { public: FlattenMapper() : OpMapper("Flatten") {} ~FlattenMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.cc index d8d23fdac2..f4594afffc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.cc @@ -20,7 +20,6 @@ #include #include #include "common/fetch_content.h" -#include "ops/op_utils.h" #include "common/op_enum.h" #include "ops/gather.h" #include "op/gather_operator.h" @@ -30,13 +29,13 @@ namespace dpico { namespace { const size_t kOfflineArgSize2 = 2; } // namespace -STATUS GatherMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS GatherMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto gather_prim = utils::cast>(prim); + auto gather_prim = api::utils::cast>(prim); MS_ASSERT(gather_prim != nullptr); auto gather_operator = std::make_unique(); @@ -61,7 +60,7 @@ STATUS GatherMapper::Map(const CNodePtr &cnode, std::vector *ba auto data = reinterpret_cast(data_info.data_.data()); gather_operator->SetAxis(*data); } else if (gather_prim->GetAttr(ops::kAxis) != nullptr) { - gather_operator->SetAxis(GetValue(gather_prim->GetAttr(ops::kAxis))); + gather_operator->SetAxis(static_cast(api::GetValue(gather_prim->GetAttr(ops::kAxis)))); } else { MS_LOG(ERROR) << "null param"; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.h index edaedc61b8..04d14603e8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/gather_mapper.h @@ -28,8 +28,8 @@ class GatherMapper : public OpMapper { public: GatherMapper() : OpMapper("Gather") {} ~GatherMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.cc index b4bd2b5a55..69947ed078 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.cc @@ -25,8 +25,8 @@ namespace mindspore { namespace dpico { -STATUS GruMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS GruMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { MS_CHECK_TRUE_MSG(base_operators != nullptr, RET_ERROR, "base_operators is nullptr."); auto gru_operator = std::make_unique(); MS_CHECK_TRUE_MSG(gru_operator != nullptr, RET_ERROR, "gru_operator is nullptr."); @@ -37,50 +37,50 @@ STATUS GruMapper::Map(const CNodePtr &cnode, std::vector *base_ } if (prim->GetAttr(kNumOutput) != nullptr) { - gru_operator->SetRecurrentNumOutput(GetValue(prim->GetAttr(kNumOutput))); + gru_operator->SetRecurrentNumOutput(static_cast(api::GetValue(prim->GetAttr(kNumOutput)))); } if (prim->GetAttr(kExposeHidden) != nullptr) { - gru_operator->SetRecurrentExposeHidden(GetValue(prim->GetAttr(kExposeHidden))); + gru_operator->SetRecurrentExposeHidden(api::GetValue(prim->GetAttr(kExposeHidden))); } if (prim->GetAttr(kHasSplitBiasFlag) != nullptr) { - gru_operator->SetHasSplitBiasFlag(GetValue(prim->GetAttr(kHasSplitBiasFlag))); + gru_operator->SetHasSplitBiasFlag(api::GetValue(prim->GetAttr(kHasSplitBiasFlag))); } if (prim->GetAttr(kHasSplitHWeightFlag) != nullptr) { - gru_operator->SetHasSplitHWeightFlag(GetValue(prim->GetAttr(kHasSplitHWeightFlag))); + gru_operator->SetHasSplitHWeightFlag(api::GetValue(prim->GetAttr(kHasSplitHWeightFlag))); } if (prim->GetAttr(kGruWeightOrderZrhFlag) != nullptr) { - gru_operator->SetGruWeightOrderZrhFlag(GetValue(prim->GetAttr(kGruWeightOrderZrhFlag))); + gru_operator->SetGruWeightOrderZrhFlag(api::GetValue(prim->GetAttr(kGruWeightOrderZrhFlag))); } if (prim->GetAttr(kOnnxModeOutFlag) != nullptr) { - gru_operator->SetOnnxModeOutFlag(GetValue(prim->GetAttr(kOnnxModeOutFlag))); + gru_operator->SetOnnxModeOutFlag(api::GetValue(prim->GetAttr(kOnnxModeOutFlag))); } if (prim->GetAttr(kOutputLastFrameFlag) != nullptr) { - gru_operator->SetOutputLastFrameFlag(GetValue(prim->GetAttr(kOutputLastFrameFlag))); + gru_operator->SetOutputLastFrameFlag(api::GetValue(prim->GetAttr(kOutputLastFrameFlag))); } if (prim->GetAttr(kKeepDirectionDimFlag) != nullptr) { - gru_operator->SetKeepDirectionDimFlag(GetValue(prim->GetAttr(kKeepDirectionDimFlag))); + gru_operator->SetKeepDirectionDimFlag(api::GetValue(prim->GetAttr(kKeepDirectionDimFlag))); } if (prim->GetAttr(kInitialHOnlineFlag) != nullptr) { - gru_operator->SetInitialHOnlineFlag(GetValue(prim->GetAttr(kInitialHOnlineFlag))); + gru_operator->SetInitialHOnlineFlag(api::GetValue(prim->GetAttr(kInitialHOnlineFlag))); } if (prim->GetAttr(kUseDefaultInitialHFlag) != nullptr) { - gru_operator->SetUseDefaultInitialHFlag(GetValue(prim->GetAttr(kUseDefaultInitialHFlag))); + gru_operator->SetUseDefaultInitialHFlag(api::GetValue(prim->GetAttr(kUseDefaultInitialHFlag))); } if (prim->GetAttr(kActivateType) != nullptr) { - gru_operator->PushActivateType(GetValue(prim->GetAttr(kActivateType))); + gru_operator->PushActivateType(api::GetValue(prim->GetAttr(kActivateType))); } if (prim->GetAttr(kActivateAlpha) != nullptr) { - gru_operator->PushActivateAlpha(GetValue(prim->GetAttr(kActivateAlpha))); + gru_operator->PushActivateAlpha(api::GetValue(prim->GetAttr(kActivateAlpha))); } if (prim->GetAttr(kActivateBeta) != nullptr) { - gru_operator->PushActivateBeta(GetValue(prim->GetAttr(kActivateBeta))); + gru_operator->PushActivateBeta(api::GetValue(prim->GetAttr(kActivateBeta))); } if (prim->GetAttr(kAfClip) != nullptr) { - gru_operator->SetAfClip(GetValue(prim->GetAttr(kAfClip))); + gru_operator->SetAfClip(api::GetValue(prim->GetAttr(kAfClip))); } if (prim->GetAttr(kRecurrentDirection) != nullptr) { - gru_operator->SetRecurrentDirection( - static_cast(GetValue(prim->GetAttr(kRecurrentDirection)))); + gru_operator->SetRecurrentDirection(static_cast( + static_cast(api::GetValue(prim->GetAttr(kRecurrentDirection))))); } if (SetRecurrentDataInfo(cnode, gru_operator.get()) != RET_OK) { MS_LOG(ERROR) << "set gru data info failed."; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.h index 084f8ed57e..4f05beafa6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/gru_mapper.h @@ -27,8 +27,8 @@ class GruMapper : public OpMapper { public: GruMapper() : OpMapper("Gru") {} ~GruMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.cc index f7e1505cf3..13e768232b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.cc @@ -19,20 +19,19 @@ #include #include #include "common/op_attr.h" -#include "ops/op_utils.h" #include "common/anf_util.h" #include "ops/resize.h" #include "op/interp_operator.h" namespace mindspore { namespace dpico { -STATUS InterpMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS InterpMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto interp_prim = utils::cast>(prim); + auto interp_prim = api::utils::cast>(prim); MS_ASSERT(interp_prim != nullptr); auto interp_operator = std::make_unique(); @@ -54,16 +53,16 @@ STATUS InterpMapper::Map(const CNodePtr &cnode, std::vector *ba interp_operator->SetInterpWidth(static_cast(interp_prim->get_new_width())); } if (prim->GetAttr(dpico::kZoomFactor) != nullptr) { - interp_operator->SetInterpZoom(GetValue(prim->GetAttr(dpico::kZoomFactor))); + interp_operator->SetInterpZoom(static_cast(api::GetValue(prim->GetAttr(dpico::kZoomFactor)))); } if (prim->GetAttr(dpico::kShrinkFactor) != nullptr) { - interp_operator->SetInterpShrink(GetValue(prim->GetAttr(dpico::kShrinkFactor))); + interp_operator->SetInterpShrink(static_cast(api::GetValue(prim->GetAttr(dpico::kShrinkFactor)))); } if (prim->GetAttr(dpico::kPadBeg) != nullptr) { - interp_operator->SetInterpPadBeg(GetValue(prim->GetAttr(dpico::kPadBeg))); + interp_operator->SetInterpPadBeg(static_cast(api::GetValue(prim->GetAttr(dpico::kPadBeg)))); } if (prim->GetAttr(dpico::kPadEnd) != nullptr) { - interp_operator->SetInterpPadEnd(GetValue(prim->GetAttr(dpico::kPadEnd))); + interp_operator->SetInterpPadEnd(static_cast(api::GetValue(prim->GetAttr(dpico::kPadEnd)))); } base_operators->push_back(std::move(interp_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.h index 6e62a17025..9ef6ab38dc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/interp_mapper.h @@ -28,8 +28,8 @@ class InterpMapper : public OpMapper { public: InterpMapper() : OpMapper("Interp") {} ~InterpMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.cc index b0f6b1db1d..7635ebfeb1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.cc @@ -24,13 +24,13 @@ namespace mindspore { namespace dpico { -STATUS LogMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS LogMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto log_prim = utils::cast>(prim); + auto log_prim = api::utils::cast>(prim); MS_ASSERT(log_prim != nullptr); auto log_operator = std::make_unique(); @@ -46,13 +46,13 @@ STATUS LogMapper::Map(const CNodePtr &cnode, std::vector *base_ log_operator->SetOpType(mapper::OpType::LOG); if (prim->GetAttr(ops::kBase) != nullptr) { - log_operator->SetLogBase(GetValue(prim->GetAttr(ops::kBase))); + log_operator->SetLogBase(api::GetValue(prim->GetAttr(ops::kBase))); } if (prim->GetAttr(ops::kScale) != nullptr) { - log_operator->SetLogScale(GetValue(prim->GetAttr(ops::kScale))); + log_operator->SetLogScale(api::GetValue(prim->GetAttr(ops::kScale))); } if (prim->GetAttr(ops::kShift) != nullptr) { - log_operator->SetLogShift(GetValue(prim->GetAttr(ops::kShift))); + log_operator->SetLogShift(api::GetValue(prim->GetAttr(ops::kShift))); } if (PushOfflineArgs(cnode, log_operator.get(), 1) != RET_OK) { MS_LOG(ERROR) << "push offline args failed. " << cnode->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.h index 61d180407d..99dd0ca67e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/log_mapper.h @@ -28,8 +28,8 @@ class LogMapper : public OpMapper { public: LogMapper() : OpMapper("Log") {} ~LogMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.cc index f73a521e3e..41eafa2bf3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.cc @@ -25,13 +25,13 @@ namespace mindspore { namespace dpico { -STATUS LrnMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS LrnMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto lrn_prim = utils::cast>(prim); + auto lrn_prim = api::utils::cast>(prim); MS_ASSERT(lrn_prim != nullptr); auto lrn_operator = std::make_unique(); @@ -51,7 +51,7 @@ STATUS LrnMapper::Map(const CNodePtr &cnode, std::vector *base_ lrn_operator->SetLrnAlpha(lrn_prim->get_alpha() * local_size); lrn_operator->SetLrnBeta(lrn_prim->get_beta()); if (lrn_prim->GetAttr(kLrnK) != nullptr) { - lrn_operator->SetLrnK(GetValue(lrn_prim->GetAttr(kLrnK))); + lrn_operator->SetLrnK(api::GetValue(lrn_prim->GetAttr(kLrnK))); } if (PushOfflineArgs(cnode, lrn_operator.get(), 1) != RET_OK) { MS_LOG(ERROR) << "push offline args failed. " << cnode->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.h index a6e769d9f5..6b0d6bff45 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/lrn_mapper.h @@ -28,8 +28,8 @@ class LrnMapper : public OpMapper { public: LrnMapper() : OpMapper("Lrn") {} ~LrnMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.cc index 171e219e6d..f68f4e05a6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS LstmMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS LstmMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -42,37 +42,37 @@ STATUS LstmMapper::Map(const CNodePtr &cnode, std::vector *base } if (prim->GetAttr(kNumOutput) != nullptr) { - lstm_operator->SetRecurrentNumOutput(GetValue(prim->GetAttr(kNumOutput))); + lstm_operator->SetRecurrentNumOutput(static_cast(api::GetValue(prim->GetAttr(kNumOutput)))); } if (prim->GetAttr(kExposeHidden) != nullptr) { - lstm_operator->SetRecurrentExposeHidden(GetValue(prim->GetAttr(kExposeHidden))); + lstm_operator->SetRecurrentExposeHidden(api::GetValue(prim->GetAttr(kExposeHidden))); } if (prim->GetAttr(kOutputLastFrameFlag) != nullptr) { - lstm_operator->SetOutputLastFrameFlag(GetValue(prim->GetAttr(kOutputLastFrameFlag))); + lstm_operator->SetOutputLastFrameFlag(api::GetValue(prim->GetAttr(kOutputLastFrameFlag))); } if (prim->GetAttr(kInitialHOnlineFlag) != nullptr) { - lstm_operator->SetInitialHOnlineFlag(GetValue(prim->GetAttr(kInitialHOnlineFlag))); + lstm_operator->SetInitialHOnlineFlag(api::GetValue(prim->GetAttr(kInitialHOnlineFlag))); } if (prim->GetAttr(kUseDefaultInitialHFlag) != nullptr) { - lstm_operator->SetUseDefaultInitialHFlag(GetValue(prim->GetAttr(kUseDefaultInitialHFlag))); + lstm_operator->SetUseDefaultInitialHFlag(api::GetValue(prim->GetAttr(kUseDefaultInitialHFlag))); } if (prim->GetAttr(kInitialCOnlineFlag) != nullptr) { - lstm_operator->SetInitialCOnlineFlag(GetValue(prim->GetAttr(kInitialCOnlineFlag))); + lstm_operator->SetInitialCOnlineFlag(api::GetValue(prim->GetAttr(kInitialCOnlineFlag))); } if (prim->GetAttr(kUseDefaultInitialCFlag) != nullptr) { - lstm_operator->SetUseDefaultInitialCFlag(GetValue(prim->GetAttr(kUseDefaultInitialCFlag))); + lstm_operator->SetUseDefaultInitialCFlag(api::GetValue(prim->GetAttr(kUseDefaultInitialCFlag))); } if (prim->GetAttr(kKeepDirectionDimFlag) != nullptr) { - lstm_operator->SetKeepDirectionDimFlag(GetValue(prim->GetAttr(kKeepDirectionDimFlag))); + lstm_operator->SetKeepDirectionDimFlag(api::GetValue(prim->GetAttr(kKeepDirectionDimFlag))); } if (prim->GetAttr(kPeepHoleFlag) != nullptr) { - lstm_operator->SetPeepholeFlag(GetValue(prim->GetAttr(kPeepHoleFlag))); + lstm_operator->SetPeepholeFlag(api::GetValue(prim->GetAttr(kPeepHoleFlag))); } if (prim->GetAttr(kLstmWeightOrderIofcFlag) != nullptr) { - lstm_operator->SetLstmWeightOrderIofcFlag(GetValue(prim->GetAttr(kLstmWeightOrderIofcFlag))); + lstm_operator->SetLstmWeightOrderIofcFlag(api::GetValue(prim->GetAttr(kLstmWeightOrderIofcFlag))); } if (prim->GetAttr(kSequenceLensOnlineFlag) != nullptr) { - lstm_operator->SetSequenceLensOnlineFlag(GetValue(prim->GetAttr(kSequenceLensOnlineFlag))); + lstm_operator->SetSequenceLensOnlineFlag(api::GetValue(prim->GetAttr(kSequenceLensOnlineFlag))); } if (SetRecurrentDataInfo(cnode, lstm_operator.get()) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.h index 9c9ec07033..b48555cdbe 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/lstm_mapper.h @@ -28,8 +28,8 @@ class LstmMapper : public OpMapper { public: LstmMapper() : OpMapper("Lstm") {} ~LstmMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.cc index 945d96e806..cf795bdb96 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.cc @@ -26,8 +26,8 @@ namespace mindspore { namespace dpico { namespace { -STATUS DoMaxtixOperatorMap(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS DoMaxtixOperatorMap(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { auto matrix_operator = std::make_unique(); if (matrix_operator == nullptr) { MS_LOG(ERROR) << "matrix_operator is nullptr."; @@ -38,19 +38,19 @@ STATUS DoMaxtixOperatorMap(const CNodePtr &cnode, std::vector * return RET_ERROR; } if (prim->GetAttr(kDim1) != nullptr) { - matrix_operator->SetMatMulDim1(GetValue(prim->GetAttr(kDim1))); + matrix_operator->SetMatMulDim1(static_cast(api::GetValue(prim->GetAttr(kDim1)))); } else { MS_LOG(ERROR) << kDim1 << " attr is missed. " << cnode->fullname_with_scope(); return RET_ERROR; } if (prim->GetAttr(kDim2) != nullptr) { - matrix_operator->SetMatMulDim2(GetValue(prim->GetAttr(kDim2))); + matrix_operator->SetMatMulDim2(static_cast(api::GetValue(prim->GetAttr(kDim2)))); } else { MS_LOG(ERROR) << kDim2 << " attr is missed. " << cnode->fullname_with_scope(); return RET_ERROR; } if (prim->GetAttr(kDim3) != nullptr) { - matrix_operator->SetMatMulDim3(GetValue(prim->GetAttr(kDim3))); + matrix_operator->SetMatMulDim3(static_cast(api::GetValue(prim->GetAttr(kDim3)))); } else { MS_LOG(ERROR) << kDim3 << " attr is missed. " << cnode->fullname_with_scope(); return RET_ERROR; @@ -59,14 +59,14 @@ STATUS DoMaxtixOperatorMap(const CNodePtr &cnode, std::vector * return RET_OK; } } // namespace -STATUS MatMulMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS MatMulMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } if (prim->GetAttr(kOperatorType) != nullptr) { - auto op_type = GetValue(prim->GetAttr(kOperatorType)); + auto op_type = api::GetValue(prim->GetAttr(kOperatorType)); if (op_type == "FullConnection") { OpMapperRegistry::GetInstance()->GetOpMapper(op_type)->Map(cnode, base_operators, prim, output_cnodes); } else if (op_type == "Matrix") { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.h index 81d3f680c7..d5552526c4 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/mat_mul_mapper.h @@ -28,8 +28,8 @@ class MatMulMapper : public OpMapper { public: MatMulMapper() : OpMapper("MatMul") {} ~MatMulMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.cc index 8be7e29ca4..832213ac9e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.cc @@ -18,15 +18,14 @@ #include #include #include -#include "ops/op_utils.h" #include "common/anf_util.h" #include "common/op_attr.h" #include "op/mvn_operator.h" namespace mindspore { namespace dpico { -STATUS MvnMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS MvnMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -45,18 +44,18 @@ STATUS MvnMapper::Map(const CNodePtr &cnode, std::vector *base_ mvn_operator->SetOpType(mapper::OpType::MVN); if (prim->GetAttr(ops::kEps) != nullptr) { - mvn_operator->SetMVNEps(GetValue(prim->GetAttr(ops::kEps))); + mvn_operator->SetMVNEps(api::GetValue(prim->GetAttr(ops::kEps))); } if (prim->GetAttr(dpico::kAcrossChannels) != nullptr) { - mvn_operator->SetMVNAcrossChannels(GetValue(prim->GetAttr(dpico::kAcrossChannels))); + mvn_operator->SetMVNAcrossChannels(api::GetValue(prim->GetAttr(dpico::kAcrossChannels))); } if (prim->GetAttr(dpico::kNormalizeVariance) != nullptr) { - mvn_operator->SetMVNNormalizeVariance(GetValue(prim->GetAttr(dpico::kNormalizeVariance))); + mvn_operator->SetMVNNormalizeVariance(api::GetValue(prim->GetAttr(dpico::kNormalizeVariance))); } if (prim->GetAttr(ops::kAxes) != nullptr) { - auto axes = GetValue>(prim->GetAttr(ops::kAxes)); + auto axes = api::GetValue>(prim->GetAttr(ops::kAxes)); for (auto axis : axes) { - mvn_operator->PushMVNAxes(axis); + mvn_operator->PushMVNAxes(static_cast(axis)); } } if (PushOfflineArgs(cnode, mvn_operator.get(), 1) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.h index 2599c7ffc5..153fdd85c2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/mvn_mapper.h @@ -28,8 +28,8 @@ class MvnMapper : public OpMapper { public: MvnMapper() : OpMapper("Mvn") {} ~MvnMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.cc index b80177eb0b..6b19b80566 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.cc @@ -22,8 +22,8 @@ namespace mindspore { namespace dpico { -STATUS NopMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS NopMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.h index 117fef2063..dc4e8bd3d2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/nop_mapper.h @@ -28,8 +28,8 @@ class NopMapper : public OpMapper { public: NopMapper() : OpMapper("Nop") {} ~NopMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.cc index b5251a8c70..d7a74d2ef9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.cc @@ -18,33 +18,33 @@ #include #include #include +#include "mindapi/ir/tensor.h" #include "common/op_attr.h" -#include "ops/op_utils.h" #include "op/normalize_operator.h" namespace mindspore { namespace dpico { namespace { -STATUS SetNormalizeDataInfo(const CNodePtr &cnode, mapper::NormalizeOperator *normalize_operator) { +STATUS SetNormalizeDataInfo(const api::CNodePtr &cnode, mapper::NormalizeOperator *normalize_operator) { if (normalize_operator == nullptr) { MS_LOG(ERROR) << "normalize_operator is nullptr."; return RET_ERROR; } for (size_t i = 1; i < cnode->inputs().size(); i++) { - AnfNodePtr input_node = cnode->input(i); - if (utils::isa(input_node)) { + api::AnfNodePtr input_node = cnode->input(i); + if (api::utils::isa(input_node)) { MS_LOG(INFO) << "cnode don't have blobs"; continue; } - if (utils::isa(input_node)) { - auto input_param_node = input_node->cast(); + if (api::utils::isa(input_node)) { + auto input_param_node = input_node->cast(); if (!input_param_node->has_default()) { MS_LOG(INFO) << "graph input don't have blobs"; continue; } - auto tensor_info = std::dynamic_pointer_cast(input_param_node->default_param()); + auto tensor_info = input_param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto raw_datas = static_cast(tensor_info->data_c()); + auto raw_datas = static_cast(tensor_info->data()); auto elem_count = tensor_info->DataSize(); normalize_operator->SetNormScaleVec(std::vector(raw_datas, raw_datas + elem_count)); } else { @@ -56,8 +56,8 @@ STATUS SetNormalizeDataInfo(const CNodePtr &cnode, mapper::NormalizeOperator *no return RET_OK; } } // namespace -STATUS NormalizeMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS NormalizeMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -76,16 +76,16 @@ STATUS NormalizeMapper::Map(const CNodePtr &cnode, std::vector normalize_operator->SetOpType(mapper::OpType::NORMALIZE); if (prim->GetAttr(dpico::kAcrossSpatial) != nullptr) { - normalize_operator->SetNormAcrossSpatial(GetValue(prim->GetAttr(dpico::kAcrossSpatial))); + normalize_operator->SetNormAcrossSpatial(api::GetValue(prim->GetAttr(dpico::kAcrossSpatial))); } if (prim->GetAttr(dpico::kChannelShared) != nullptr) { - normalize_operator->SetNormChannelShared(GetValue(prim->GetAttr(dpico::kChannelShared))); + normalize_operator->SetNormChannelShared(api::GetValue(prim->GetAttr(dpico::kChannelShared))); } if (prim->GetAttr(dpico::kSqrtA) != nullptr) { - normalize_operator->SetNormAlpha(GetValue(prim->GetAttr(dpico::kSqrtA))); + normalize_operator->SetNormAlpha(api::GetValue(prim->GetAttr(dpico::kSqrtA))); } if (prim->GetAttr(ops::kEps) != nullptr) { - normalize_operator->SetNormEps(GetValue(prim->GetAttr(ops::kEps))); + normalize_operator->SetNormEps(api::GetValue(prim->GetAttr(ops::kEps))); } if (SetNormalizeDataInfo(cnode, normalize_operator.get()) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.h index 96a7f3b5fb..2526c85df4 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/normalize_mapper.h @@ -28,8 +28,8 @@ class NormalizeMapper : public OpMapper { public: NormalizeMapper() : OpMapper("Normalize") {} ~NormalizeMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.cc index b2fe444a21..578a4c69e0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.cc @@ -15,17 +15,19 @@ */ #include "mapper/op_mapper.h" -#include #include +#include +#include "ops/tuple_get_item.h" #include "common/op_attr.h" #include "common/op_enum.h" #include "common/anf_util.h" #include "common/string_util.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace dpico { namespace { -STATUS SetOpInputs(const CNodePtr &cnode, mapper::BaseOperator *base_operator) { +STATUS SetOpInputs(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator) { if (base_operator == nullptr) { MS_LOG(ERROR) << "base_operator is nullptr."; return RET_ERROR; @@ -34,20 +36,20 @@ STATUS SetOpInputs(const CNodePtr &cnode, mapper::BaseOperator *base_operator) { for (size_t i = 1; i < cnode->inputs().size(); i++) { auto input_anode = cnode->input(i); MS_ASSERT(input_anode != nullptr); - if (utils::isa(input_anode)) { - auto param_node = input_anode->cast(); + if (api::utils::isa(input_anode)) { + auto param_node = input_anode->cast(); if (param_node != nullptr && !param_node->has_default()) { // graph input input_names.emplace_back(input_anode->fullname_with_scope()); } - } else if (utils::isa(input_anode)) { - auto input_cnode = input_anode->cast(); + } else if (api::utils::isa(input_anode)) { + auto input_cnode = input_anode->cast(); if (input_cnode == nullptr) { MS_LOG(ERROR) << "input node must be cnode."; return RET_ERROR; } auto node_name = input_cnode->fullname_with_scope(); if (input_cnode->GetAttr(kOutputsNames) != nullptr) { - auto output_names = GetValue>(input_cnode->GetAttr(kOutputsNames)); + auto output_names = api::GetValue>(input_cnode->GetAttr(kOutputsNames)); if (output_names.size() == 1) { node_name = output_names.at(0); } @@ -59,11 +61,12 @@ STATUS SetOpInputs(const CNodePtr &cnode, mapper::BaseOperator *base_operator) { return RET_OK; } -STATUS FillMultiOutOpOutputs(const CNodePtr &cnode, mapper::BaseOperator *base_operator, - const CNodePtrList &output_cnodes) { +STATUS FillMultiOutOpOutputs(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator, + const api::CNodePtrList &output_cnodes) { MS_ASSERT(base_operator != nullptr); - if (std::any_of(output_cnodes.begin(), output_cnodes.end(), - [](const CNodePtr cnode) { return !CheckPrimitiveType(cnode, prim::kPrimTupleGetItem); })) { + if (std::any_of(output_cnodes.begin(), output_cnodes.end(), [](const api::CNodePtr &cnode) { + return !CheckPrimitiveType(cnode, api::MakeShared()); + })) { MS_LOG(ERROR) << "multi-out op must be connected with tuple-get-item node."; return RET_ERROR; } @@ -72,11 +75,11 @@ STATUS FillMultiOutOpOutputs(const CNodePtr &cnode, mapper::BaseOperator *base_o MS_LOG(ERROR) << "each node's abstract must be not a nullptr."; return RET_ERROR; } - if (!abstract->isa()) { + if (!abstract->isa()) { MS_LOG(ERROR) << "multi-out op's abstract must be a tuple."; return RET_ERROR; } - auto abstract_tuple = abstract->cast(); + auto abstract_tuple = abstract->cast(); MS_ASSERT(abstract_tuple != nullptr); auto output_num = abstract_tuple->elements().size(); std::vector output_names; @@ -84,14 +87,14 @@ STATUS FillMultiOutOpOutputs(const CNodePtr &cnode, mapper::BaseOperator *base_o for (size_t i = 0; i < output_num; ++i) { output_names.emplace_back(cnode->fullname_with_scope() + "_unused_" + std::to_string(i)); } - for (auto output_cnode : output_cnodes) { + for (const auto &output_cnode : output_cnodes) { if (output_cnode->size() != kInputIndex3) { MS_LOG(ERROR) << "tuple-get_item's inputs size must be 3."; return RET_ERROR; } auto index_node = output_cnode->input(kInputIndex2); MS_CHECK_TRUE_MSG(index_node != nullptr, RET_ERROR, "node is nullptr."); - auto value_ptr = GetValueNode(index_node); + auto value_ptr = api::GetValueNode(index_node); MS_CHECK_TRUE_MSG(value_ptr != nullptr, RET_ERROR, "tuple_get_item's second input must be a value."); auto num_str = value_ptr->ToString(); MS_CHECK_TRUE_MSG(IsValidUnsignedNum(num_str), RET_ERROR, "tuple_get_item's second input must be an unsigned int"); @@ -104,14 +107,17 @@ STATUS FillMultiOutOpOutputs(const CNodePtr &cnode, mapper::BaseOperator *base_o return RET_OK; } -STATUS SetOpOutputs(const CNodePtr &cnode, mapper::BaseOperator *base_operator, const CNodePtrList &output_cnodes) { +STATUS SetOpOutputs(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator, + const api::CNodePtrList &output_cnodes) { if (cnode == nullptr || base_operator == nullptr || - std::any_of(output_cnodes.begin(), output_cnodes.end(), [](const CNodePtr cnode) { return cnode == nullptr; })) { + std::any_of(output_cnodes.begin(), output_cnodes.end(), + [](const api::CNodePtr &cnode) { return cnode == nullptr; })) { MS_LOG(ERROR) << "the function exist that input parameter is a nullptr."; return RET_ERROR; } - if (std::all_of(output_cnodes.begin(), output_cnodes.end(), - [](const CNodePtr cnode) { return !CheckPrimitiveType(cnode, prim::kPrimTupleGetItem); })) { + if (std::all_of(output_cnodes.begin(), output_cnodes.end(), [](const api::CNodePtr &cnode) { + return !CheckPrimitiveType(cnode, api::MakeShared()); + })) { // single output op std::vector output_names; output_names.emplace_back(cnode->fullname_with_scope()); @@ -128,7 +134,8 @@ STATUS SetOpOutputs(const CNodePtr &cnode, mapper::BaseOperator *base_operator, } } // namespace -STATUS SetCommonAttr(const CNodePtr &cnode, mapper::BaseOperator *base_operator, const CNodePtrList &output_cnodes) { +STATUS SetCommonAttr(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator, + const api::CNodePtrList &output_cnodes) { if (base_operator == nullptr) { MS_LOG(ERROR) << "base operator is nullptr."; return RET_ERROR; @@ -146,7 +153,7 @@ STATUS SetCommonAttr(const CNodePtr &cnode, mapper::BaseOperator *base_operator, return RET_OK; } -STATUS SetConvFcDataInfo(const CNodePtr &cnode, mapper::BaseOperator *base_operator) { +STATUS SetConvFcDataInfo(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator) { if (base_operator == nullptr) { MS_LOG(ERROR) << "base_operator is nullptr."; return RET_ERROR; @@ -154,13 +161,13 @@ STATUS SetConvFcDataInfo(const CNodePtr &cnode, mapper::BaseOperator *base_opera for (size_t i = 2; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); MS_ASSERT(input_node != nullptr); - auto param_node = input_node->cast(); + auto param_node = input_node->cast(); if (param_node == nullptr || !param_node->has_default()) { continue; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); if (i == kInputIndex2) { base_operator->SetWeightDataPtr(data); @@ -181,26 +188,26 @@ STATUS SetConvFcDataInfo(const CNodePtr &cnode, mapper::BaseOperator *base_opera return RET_OK; } -STATUS SetRecurrentDataInfo(const CNodePtr &cnode, mapper::RecurrentOperator *recurrent_operator) { +STATUS SetRecurrentDataInfo(const api::CNodePtr &cnode, mapper::RecurrentOperator *recurrent_operator) { if (recurrent_operator == nullptr) { MS_LOG(ERROR) << "recurrent_operator is nullptr."; return RET_ERROR; } for (size_t i = 1; i < cnode->inputs().size(); i++) { - AnfNodePtr input_node = cnode->input(i); - if (utils::isa(input_node)) { + auto input_node = cnode->input(i); + if (api::utils::isa(input_node)) { MS_LOG(INFO) << "cnode don't have blobs"; continue; } - if (utils::isa(input_node)) { - auto input_param_node = input_node->cast(); + if (api::utils::isa(input_node)) { + auto input_param_node = input_node->cast(); if (!input_param_node->has_default()) { MS_LOG(INFO) << "graph input don't have blobs"; continue; } - auto tensor_info = std::dynamic_pointer_cast(input_param_node->default_param()); + auto tensor_info = input_param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto raw_datas = static_cast(tensor_info->data_c()); + auto raw_datas = static_cast(tensor_info->data()); auto elem_count = tensor_info->DataSize(); auto weight_data = new (std::nothrow) float[tensor_info->DataSize()]; if (weight_data == nullptr) { @@ -223,7 +230,7 @@ STATUS SetRecurrentDataInfo(const CNodePtr &cnode, mapper::RecurrentOperator *re } return RET_OK; } -STATUS PushOfflineArgs(const CNodePtr &cnode, mapper::BaseOperator *base_operator, size_t offline_args_size) { +STATUS PushOfflineArgs(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator, size_t offline_args_size) { if (base_operator == nullptr) { MS_LOG(ERROR) << "base_operator is nullptr."; return RET_ERROR; @@ -238,29 +245,29 @@ STATUS PushOfflineArgs(const CNodePtr &cnode, mapper::BaseOperator *base_operato std::vector, std::vector>> offline_args; bool has_offline_args = false; for (size_t i = 1; i < inputs_size; i++) { - AnfNodePtr input_node = cnode->input(i); - if (utils::isa(input_node)) { + auto input_node = cnode->input(i); + if (api::utils::isa(input_node)) { MS_LOG(INFO) << "cnode don't have blobs"; offline_args.emplace_back(); continue; } - if (utils::isa(input_node)) { - auto input_param_node = input_node->cast(); + if (api::utils::isa(input_node)) { + auto input_param_node = input_node->cast(); if (!input_param_node->has_default()) { MS_LOG(INFO) << "graph input don't have blobs"; offline_args.emplace_back(); continue; } - auto tensor_info = std::dynamic_pointer_cast(input_param_node->default_param()); + auto tensor_info = input_param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { has_offline_args = true; std::vector offline_data; auto elem_count = tensor_info->DataSize(); if (tensor_info->data_type() == kNumberTypeInt32 || tensor_info->data_type() == kNumberTypeInt) { - auto raw_datas = static_cast(tensor_info->data_c()); + auto raw_datas = static_cast(tensor_info->data()); offline_data = std::vector(raw_datas, raw_datas + elem_count); } else if (tensor_info->data_type() == kNumberTypeFloat32 || tensor_info->data_type() == kNumberTypeFloat) { - auto raw_datas = static_cast(tensor_info->data_c()); + auto raw_datas = static_cast(tensor_info->data()); offline_data = std::vector(raw_datas, raw_datas + elem_count); } else { MS_LOG(ERROR) << "unsupported param type. " << tensor_info->data_type(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.h index ee99662570..333490baf2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/op_mapper.h @@ -24,34 +24,37 @@ #include #include "common/check_base.h" #include "common/fetch_content.h" -#include "base/base.h" -#include "ir/anf.h" +#include "mindapi/base/base.h" +#include "mindapi/ir/anf.h" #include "include/errorcode.h" #include "op/base_operator.h" #include "op/recurrent_operator.h" +#include "mindapi/base/logging.h" +#include "ops/op_name.h" using mindspore::lite::RET_ERROR; using mindspore::lite::RET_OK; using mindspore::lite::STATUS; -using BaseOperatorPtr = std::unique_ptr; namespace mindspore { namespace dpico { +using BaseOperatorPtr = std::unique_ptr; class OpMapper { public: explicit OpMapper(std::string node_name) : name(std::move(node_name)) {} virtual ~OpMapper() = default; - virtual STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) = 0; + virtual STATUS Map(const api::CNodePtr &node, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) = 0; private: const std::string name; }; using OpMapperPtr = std::shared_ptr; -STATUS SetCommonAttr(const CNodePtr &node, mapper::BaseOperator *base_operator, const CNodePtrList &output_cnodes); -STATUS SetConvFcDataInfo(const CNodePtr &cnode, mapper::BaseOperator *base_operator); -STATUS SetRecurrentDataInfo(const CNodePtr &cnode, mapper::RecurrentOperator *recurrent_operator); -STATUS PushOfflineArgs(const CNodePtr &cnode, mapper::BaseOperator *base_operator, size_t offline_args_size); +STATUS SetCommonAttr(const api::CNodePtr &node, mapper::BaseOperator *base_operator, + const api::CNodePtrList &output_cnodes); +STATUS SetConvFcDataInfo(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator); +STATUS SetRecurrentDataInfo(const api::CNodePtr &cnode, mapper::RecurrentOperator *recurrent_operator); +STATUS PushOfflineArgs(const api::CNodePtr &cnode, mapper::BaseOperator *base_operator, size_t offline_args_size); } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.cc index 08f0414bca..19872272bd 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS PassThroughMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS PassThroughMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -44,13 +44,16 @@ STATUS PassThroughMapper::Map(const CNodePtr &cnode, std::vectorSetOpType(mapper::OpType::PASS_THROUGH); if (prim->GetAttr(kNumOutput) != nullptr) { - pass_through_operator->SetPassThroughNumOutput(GetValue(prim->GetAttr(kNumOutput))); + pass_through_operator->SetPassThroughNumOutput( + static_cast(api::GetValue(prim->GetAttr(kNumOutput)))); } if (prim->GetAttr(kBlockHeight) != nullptr) { - pass_through_operator->SetPassThroughBlockHeight(GetValue(prim->GetAttr(kBlockHeight))); + pass_through_operator->SetPassThroughBlockHeight( + static_cast(api::GetValue(prim->GetAttr(kBlockHeight)))); } if (prim->GetAttr(kBlockWidth) != nullptr) { - pass_through_operator->SetPassThroughBlockWidth(GetValue(prim->GetAttr(kBlockWidth))); + pass_through_operator->SetPassThroughBlockWidth( + static_cast(api::GetValue(prim->GetAttr(kBlockWidth)))); } base_operators->push_back(std::move(pass_through_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.h index 128549472d..be5da77f7d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/passthrough_mapper.h @@ -28,8 +28,8 @@ class PassThroughMapper : public OpMapper { public: PassThroughMapper() : OpMapper("PassThrough") {} ~PassThroughMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.cc index acc15a6258..9f9e9e7bee 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.cc @@ -18,6 +18,7 @@ #include #include #include +#include #include "common/op_enum.h" #include "common/data_transpose_utils.h" #include "common/fetch_content.h" @@ -27,13 +28,13 @@ namespace mindspore { namespace dpico { -STATUS PermuteMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS PermuteMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto permute_prim = utils::cast>(prim); + auto permute_prim = api::utils::cast>(prim); MS_ASSERT(permute_prim != nullptr); auto permute_operator = std::make_unique(); @@ -72,7 +73,9 @@ STATUS PermuteMapper::Map(const CNodePtr &cnode, std::vector *b } perm_val = {data[0], data[1], data[kAxis2], data[kAxis3]}; } else if (permute_prim->GetAttr(kPerm) != nullptr) { - perm_val = GetValue>(permute_prim->GetAttr(kPerm)); + auto perm_vec = api::GetValue>(permute_prim->GetAttr(kPerm)); + (void)std::transform(perm_vec.begin(), perm_vec.end(), std::back_inserter(perm_val), + [](int64_t p) { return static_cast(p); }); } else { MS_LOG(ERROR) << "can't get perm value. " << cnode->fullname_with_scope(); return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.h index 879018c6ac..7ec509600b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/permute_mapper.h @@ -28,8 +28,8 @@ class PermuteMapper : public OpMapper { public: PermuteMapper() : OpMapper("Permute") {} ~PermuteMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.cc index cae176b32b..76757d6a80 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.cc @@ -18,23 +18,24 @@ #include #include #include -#include "ops/op_utils.h" #include "common/op_enum.h" #include "common/anf_util.h" #include "op/pool_operator.h" +#include "ops/fusion/max_pool_fusion.h" +#include "ops/fusion/avg_pool_fusion.h" namespace mindspore { namespace dpico { namespace { -STATUS SetPadAttr(const PrimitivePtr &prim, mapper::PoolOperator *pool_operator) { +STATUS SetPadAttr(const api::PrimitivePtr &prim, mapper::PoolOperator *pool_operator) { if (pool_operator == nullptr) { MS_LOG(ERROR) << "pool_operator is nullptr. "; return RET_ERROR; } if (prim->GetAttr(ops::kPadMode) != nullptr) { - auto pad_mode = PadMode(GetValue(prim->GetAttr(ops::kPadMode))); + auto pad_mode = PadMode(api::GetValue(prim->GetAttr(ops::kPadMode))); if (pad_mode == PadMode::PAD) { - auto pad_list = GetValue>(prim->GetAttr(ops::kPad)); + auto pad_list = api::GetValue>(prim->GetAttr(ops::kPad)); if (pad_list.size() != kDims4) { MS_LOG(ERROR) << "pad_list size is invalid. " << pad_list.size(); return RET_ERROR; @@ -63,8 +64,8 @@ STATUS SetPadAttr(const PrimitivePtr &prim, mapper::PoolOperator *pool_operator) return RET_OK; } } // namespace -STATUS PoolMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS PoolMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -80,12 +81,12 @@ STATUS PoolMapper::Map(const CNodePtr &cnode, std::vector *base MS_LOG(ERROR) << "set common attr failed. " << cnode->fullname_with_scope(); return RET_ERROR; } - if (CheckPrimitiveType(cnode, prim::kPrimAvgPoolFusion)) { + if (CheckPrimitiveType(cnode, api::MakeShared())) { pool_operator->SetOpType(mapper::OpType::POOLINGAVE); - } else if (CheckPrimitiveType(cnode, prim::kPrimMaxPoolFusion)) { + } else if (CheckPrimitiveType(cnode, api::MakeShared())) { pool_operator->SetOpType(mapper::OpType::POOLINGMAX); } else { - auto primitive = mindspore::GetValueNode(cnode->input(0)); + auto primitive = api::GetValueNode(cnode->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr. " << cnode->fullname_with_scope(); return RET_ERROR; @@ -95,7 +96,7 @@ STATUS PoolMapper::Map(const CNodePtr &cnode, std::vector *base } if (prim->GetAttr(ops::kKernelSize)) { - auto kernel_size = GetValue>(prim->GetAttr(ops::kKernelSize)); + auto kernel_size = api::GetValue>(prim->GetAttr(ops::kKernelSize)); if (kernel_size.size() != kDims2) { MS_LOG(ERROR) << "kernel_size should be 2 dims, which is " << kernel_size.size(); return RET_ERROR; @@ -105,7 +106,7 @@ STATUS PoolMapper::Map(const CNodePtr &cnode, std::vector *base } if (prim->GetAttr(ops::kStrides) != nullptr) { - auto stride = GetValue>(prim->GetAttr(ops::kStrides)); + auto stride = api::GetValue>(prim->GetAttr(ops::kStrides)); if (stride.size() != kDims2) { MS_LOG(ERROR) << "stride should be 2 dims, which is " << stride.size(); return RET_ERROR; @@ -115,12 +116,12 @@ STATUS PoolMapper::Map(const CNodePtr &cnode, std::vector *base } if (prim->GetAttr(ops::kGlobal) != nullptr) { - auto global_flag = GetValue(prim->GetAttr(ops::kGlobal)); + auto global_flag = api::GetValue(prim->GetAttr(ops::kGlobal)); pool_operator->SetGlobalPoolingFlag(global_flag); } if (prim->GetAttr(ops::kRoundMode) != nullptr) { - auto round_mode = GetValue(prim->GetAttr(ops::kRoundMode)); + auto round_mode = api::GetValue(prim->GetAttr(ops::kRoundMode)); if (round_mode == RoundMode::CEIL) { pool_operator->SetRoundMode(mapper::POOLING_ROUND_MODE_CEIL); } else if (round_mode == RoundMode::FLOOR) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.h index c7d831f7bb..237916f538 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/pool_mapper.h @@ -30,8 +30,8 @@ class PoolMapper : public OpMapper { public: PoolMapper() : OpMapper("Pool") {} ~PoolMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.cc index 1d4315d992..f857bd9690 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.cc @@ -26,13 +26,13 @@ namespace mindspore { namespace dpico { -STATUS PowerMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS PowerMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto power_prim = utils::cast>(prim); + auto power_prim = api::utils::cast>(prim); MS_ASSERT(power_prim != nullptr); auto power_operator = std::make_unique(); @@ -61,16 +61,16 @@ STATUS PowerMapper::Map(const CNodePtr &cnode, std::vector *bas } power_operator->SetPowerPower(*data); } else if (prim->GetAttr(ops::kPower) != nullptr) { - power_operator->SetPowerPower(GetValue(prim->GetAttr(ops::kPower))); + power_operator->SetPowerPower(api::GetValue(prim->GetAttr(ops::kPower))); } else { MS_LOG(ERROR) << "null param"; return RET_ERROR; } if (prim->GetAttr(ops::kScale) != nullptr) { - power_operator->SetPowerScale(GetValue(prim->GetAttr(ops::kScale))); + power_operator->SetPowerScale(api::GetValue(prim->GetAttr(ops::kScale))); } if (prim->GetAttr(ops::kShift) != nullptr) { - power_operator->SetPowerShift(GetValue(prim->GetAttr(ops::kShift))); + power_operator->SetPowerShift(api::GetValue(prim->GetAttr(ops::kShift))); } if (PushOfflineArgs(cnode, power_operator.get(), 1) != RET_OK) { MS_LOG(ERROR) << "push offline args failed. " << cnode->fullname_with_scope(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.h index b5867fa0b6..ddfd23f9b4 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/power_mapper.h @@ -28,8 +28,8 @@ class PowerMapper : public OpMapper { public: PowerMapper() : OpMapper("Power") {} ~PowerMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.cc index 74cd18877f..10bb0d309b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.cc @@ -19,7 +19,6 @@ #include #include #include "common/op_enum.h" -#include "ops/op_utils.h" #include "common/fetch_content.h" #include "common/anf_util.h" #include "ops/fusion/prelu_fusion.h" @@ -28,20 +27,21 @@ namespace mindspore { namespace dpico { namespace { -STATUS SetPReluDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper::PreluOperator *prelu_operator) { +STATUS SetPReluDataInfo(const api::CNodePtr &cnode, const api::PrimitivePtr &prim, + mapper::PreluOperator *prelu_operator) { if (prim->GetAttr(ops::kSlope) != nullptr) { - prelu_operator->SetAlphaNegVec(GetValue>(prim->GetAttr(ops::kSlope))); + prelu_operator->SetAlphaNegVec(api::GetValue>(prim->GetAttr(ops::kSlope))); } else if (cnode->inputs().size() > kInputIndex2) { auto input_anode = cnode->input(kInputIndex2); - if (utils::isa(input_anode)) { - auto input_param_node = input_anode->cast(); + if (api::utils::isa(input_anode)) { + auto input_param_node = input_anode->cast(); if (input_param_node == nullptr) { MS_LOG(ERROR) << input_param_node->fullname_with_scope() << " is nullptr."; return RET_ERROR; } - auto tensor_info = std::dynamic_pointer_cast(input_param_node->default_param()); + auto tensor_info = input_param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto raw_datas = static_cast(tensor_info->data_c()); + auto raw_datas = static_cast(tensor_info->data()); auto elem_count = tensor_info->DataSize(); prelu_operator->SetAlphaNegVec(std::vector(raw_datas, raw_datas + elem_count)); } else { @@ -53,13 +53,13 @@ STATUS SetPReluDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper: return RET_OK; } } // namespace -STATUS PReluMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS PReluMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto prelu_prim = utils::cast>(prim); + auto prelu_prim = api::utils::cast>(prim); MS_ASSERT(prelu_prim != nullptr); auto prelu_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.h index f82214a847..7980218465 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/prelu_mapper.h @@ -28,8 +28,8 @@ class PReluMapper : public OpMapper { public: PReluMapper() : OpMapper("PRelu") {} ~PReluMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.cc index 1fb9564b27..759e634145 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS PsRoiPoolMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS PsRoiPoolMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -44,13 +44,13 @@ STATUS PsRoiPoolMapper::Map(const CNodePtr &cnode, std::vector psroi_pool_operator->SetOpType(mapper::OpType::PSROI); if (prim->GetAttr(kSpatialScale) != nullptr) { - psroi_pool_operator->SetPsroiSpatialScale(GetValue(prim->GetAttr(kSpatialScale))); + psroi_pool_operator->SetPsroiSpatialScale(api::GetValue(prim->GetAttr(kSpatialScale))); } if (prim->GetAttr(kOutputDim) != nullptr) { - psroi_pool_operator->SetPsroiOutputDim(GetValue(prim->GetAttr(kOutputDim))); + psroi_pool_operator->SetPsroiOutputDim(static_cast(api::GetValue(prim->GetAttr(kOutputDim)))); } if (prim->GetAttr(kGroupSize) != nullptr) { - psroi_pool_operator->SetPsroiGroupSize(GetValue(prim->GetAttr(kGroupSize))); + psroi_pool_operator->SetPsroiGroupSize(static_cast(api::GetValue(prim->GetAttr(kGroupSize)))); } base_operators->push_back(std::move(psroi_pool_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.h index b168b979eb..bb4b41a951 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/psroi_pool_mapper.h @@ -28,8 +28,8 @@ class PsRoiPoolMapper : public OpMapper { public: PsRoiPoolMapper() : OpMapper("PsRoiPool") {} ~PsRoiPoolMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.cc index 23a1b0d5c4..3ca022e134 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.cc @@ -23,14 +23,13 @@ #include "include/registry/converter_context.h" #include "common/op_enum.h" #include "common/fetch_content.h" -#include "ops/op_utils.h" #include "ops/fusion/reduce_fusion.h" #include "op/reduction_operator.h" namespace mindspore { namespace dpico { namespace { -int GetReductionAxes(const CNodePtr &cnode, const std::shared_ptr &reduction_prim, +int GetReductionAxes(const api::CNodePtr &cnode, const api::SharedPtr &reduction_prim, DataInfo *data_info, std::set *axes) { if (data_info == nullptr || axes == nullptr) { MS_LOG(ERROR) << "input arg is nullptr." << cnode->fullname_with_scope(); @@ -51,19 +50,21 @@ int GetReductionAxes(const CNodePtr &cnode, const std::shared_ptr(value); }); } else if (reduction_prim->GetAttr(ops::kAxes) != nullptr) { - auto axes_vec = GetValue>(reduction_prim->GetAttr(ops::kAxes)); + auto axes_vec = api::GetValue>(reduction_prim->GetAttr(ops::kAxes)); + (void)std::transform(axes_vec.begin(), axes_vec.end(), std::inserter(*axes, (*axes).begin()), + [](int64_t axis) { return static_cast(axis); }); *axes = std::set(axes_vec.begin(), axes_vec.end()); } return RET_OK; } } // namespace -STATUS ReductionMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS ReductionMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto reduction_prim = utils::cast>(prim); + auto reduction_prim = api::utils::cast>(prim); MS_ASSERT(reduction_prim != nullptr); auto reduction_operator = std::make_unique(); @@ -118,14 +119,14 @@ STATUS ReductionMapper::Map(const CNodePtr &cnode, std::vector reduction_operator->SetAxis(*axes.begin()); if (reduction_prim->GetAttr(ops::kCoeff) != nullptr) { - reduction_operator->SetReductionCoeff(GetValue(reduction_prim->GetAttr(ops::kCoeff))); + reduction_operator->SetReductionCoeff(api::GetValue(reduction_prim->GetAttr(ops::kCoeff))); } if (reduction_prim->GetAttr(ops::kKeepDims) != nullptr) { reduction_operator->SetReduceKeepDims( - static_cast(GetValue(reduction_prim->GetAttr(ops::kKeepDims)))); + static_cast(api::GetValue(reduction_prim->GetAttr(ops::kKeepDims)))); } if (reduction_prim->GetAttr(ops::kFmkType) != nullptr) { - auto fmk_type = static_cast(GetValue(reduction_prim->GetAttr(ops::kFmkType))); + auto fmk_type = static_cast(api::GetValue(reduction_prim->GetAttr(ops::kFmkType))); if (fmk_type == converter::kFmkTypeCaffe) { reduction_operator->SetReduceIsFromCaffe(true); } diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.h index 8b564566dc..31e49a37fd 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/reduction_mapper.h @@ -28,8 +28,8 @@ class ReductionMapper : public OpMapper { public: ReductionMapper() : OpMapper("Reduction") {} ~ReductionMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.cc index c2dc6b11f8..c23097dfb2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.cc @@ -21,20 +21,19 @@ #include #include "common/op_attr.h" #include "common/op_enum.h" -#include "ops/op_utils.h" #include "common/fetch_content.h" #include "ops/reshape.h" #include "op/reshape_operator.h" namespace mindspore { namespace dpico { -STATUS ReshapeMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ReshapeMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto reshape_prim = utils::cast>(prim); + auto reshape_prim = api::utils::cast>(prim); MS_ASSERT(reshape_prim != nullptr); auto reshape_operator = std::make_unique(); @@ -72,7 +71,7 @@ STATUS ReshapeMapper::Map(const CNodePtr &cnode, std::vector *b [](const int32_t &value) { return static_cast(value); }); reshape_operator->SetReshapeDimVec(dims); } else if (reshape_prim->GetAttr(ops::kShape) != nullptr) { - auto shape = GetValue>(reshape_prim->GetAttr(ops::kShape)); + auto shape = api::GetValue>(reshape_prim->GetAttr(ops::kShape)); std::vector dims; (void)std::transform(shape.begin(), shape.end(), std::back_inserter(dims), [](const int64_t &value) { return static_cast(value); }); @@ -82,13 +81,13 @@ STATUS ReshapeMapper::Map(const CNodePtr &cnode, std::vector *b return RET_ERROR; } if (prim->GetAttr(ops::kAxis) != nullptr) { - auto axis = GetValue(reshape_prim->GetAttr(ops::kAxis)); + auto axis = static_cast(api::GetValue(reshape_prim->GetAttr(ops::kAxis))); reshape_operator->SetAxis(axis); } else { reshape_operator->SetAxis(0); } if (prim->GetAttr(kNumAxes) != nullptr) { - auto num_axes = GetValue(reshape_prim->GetAttr(kNumAxes)); + auto num_axes = static_cast(api::GetValue(reshape_prim->GetAttr(kNumAxes))); reshape_operator->SetReshapeNumAxes(num_axes); } if (PushOfflineArgs(cnode, reshape_operator.get(), 1) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.h index bb1ffbd31a..3e7933e815 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/reshape_mapper.h @@ -28,8 +28,8 @@ class ReshapeMapper : public OpMapper { public: ReshapeMapper() : OpMapper("Reshape") {} ~ReshapeMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.cc index 0e757f2876..b09cef6e39 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.cc @@ -22,7 +22,6 @@ #include "common/anf_util.h" #include "include/registry/converter_context.h" #include "common/op_enum.h" -#include "ops/op_utils.h" #include "ops/resize.h" #include "op/resize_operator.h" @@ -44,7 +43,7 @@ const std::unordered_map kNearestMo {mindspore::NearestMode::CEIL, mapper::NearestMode::NEAREST_CEIL}, {mindspore::NearestMode::FLOOR, mapper::NearestMode::NEAREST_FLOOR}}; -STATUS GetShapeAndFormat(const CNodePtr &cnode, const PrimitivePtr &prim, std::vector *input_shape, +STATUS GetShapeAndFormat(const api::CNodePtr &cnode, const api::PrimitivePtr &prim, std::vector *input_shape, Format *input_format) { MS_CHECK_TRUE_MSG(input_shape != nullptr, RET_ERROR, "input_shape is nullptr."); MS_CHECK_TRUE_MSG(input_format != nullptr, RET_ERROR, "input_format is nullptr."); @@ -58,14 +57,15 @@ STATUS GetShapeAndFormat(const CNodePtr &cnode, const PrimitivePtr &prim, std::v return RET_ERROR; } if (prim->GetAttr(ops::kFormat) != nullptr) { - *input_format = static_cast(GetValue(prim->GetAttr(ops::kFormat))); + *input_format = static_cast(api::GetValue(prim->GetAttr(ops::kFormat))); } else { MS_LOG(ERROR) << ops::kFormat << " attr is needed."; return RET_ERROR; } return RET_OK; } -STATUS SetResizeDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper::ResizeOperator *resize_operator) { +STATUS SetResizeDataInfo(const api::CNodePtr &cnode, const api::PrimitivePtr &prim, + mapper::ResizeOperator *resize_operator) { MS_CHECK_TRUE_MSG(resize_operator != nullptr, RET_ERROR, "resize_operator is nullptr."); if (cnode->inputs().size() != dpico::kDims3) { MS_LOG(DEBUG) << "only process two inputs. " << cnode->fullname_with_scope(); @@ -73,12 +73,12 @@ STATUS SetResizeDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper } auto input_node = cnode->input(kAxis2); MS_CHECK_TRUE_MSG(input_node != nullptr, RET_ERROR, "input_node is nullptr."); - auto param_node = input_node->cast(); + auto param_node = input_node->cast(); if (param_node == nullptr || !param_node->has_default()) { MS_LOG(ERROR) << "invalid parameter node. " << cnode->fullname_with_scope(); return RET_ERROR; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info == nullptr || tensor_info->DataSize() == 0) { MS_LOG(ERROR) << "tensor_info is invalid. " << cnode->fullname_with_scope(); return RET_ERROR; @@ -102,7 +102,7 @@ STATUS SetResizeDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper "resize param element size should be 2. " << cnode->fullname_with_scope()); if (tensor_info->data_type() == kNumberTypeInt32 || tensor_info->data_type() == kNumberTypeInt) { std::vector size_vec; - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); if (input_format == Format::NCHW) { size_vec = {static_cast(input_shape.at(0)), static_cast(input_shape.at(kNCHW_C)), *data, @@ -114,7 +114,7 @@ STATUS SetResizeDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper resize_operator->SetSizeVec(size_vec); } else if (tensor_info->data_type() == kNumberTypeFloat32 || tensor_info->data_type() == kNumberTypeFloat) { std::vector scale_vec; - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); if (input_format == Format::NCHW) { scale_vec = {1.0, 1.0, *data, *(data + 1)}; @@ -129,17 +129,17 @@ STATUS SetResizeDataInfo(const CNodePtr &cnode, const PrimitivePtr &prim, mapper return RET_OK; } } // namespace -STATUS ResizeMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ResizeMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto resize_prim = utils::cast>(prim); + auto resize_prim = api::utils::cast>(prim); MS_ASSERT(resize_prim != nullptr); if (resize_prim->GetAttr(ops::kFmkType) != nullptr) { - auto fmk_type = static_cast(GetValue(resize_prim->GetAttr(ops::kFmkType))); + auto fmk_type = static_cast(api::GetValue(resize_prim->GetAttr(ops::kFmkType))); if (fmk_type == converter::kFmkTypeCaffe) { MS_CHECK_TRUE_MSG(OpMapperRegistry::GetInstance()->GetOpMapper("Interp") != nullptr, RET_ERROR, "mapper is nullptr."); @@ -178,7 +178,8 @@ STATUS ResizeMapper::Map(const CNodePtr &cnode, std::vector *ba if (prim->GetAttr(ops::kCoordinateTransformMode) != nullptr) { auto coordinate_transform_mode = static_cast(resize_prim->get_coordinate_transform_mode()); if (kCoordinateModeMap.find(coordinate_transform_mode) == kCoordinateModeMap.end()) { - MS_LOG(ERROR) << "unsupported coordinate transform mode:" << coordinate_transform_mode << " " + MS_LOG(ERROR) << "unsupported coordinate transform mode:" + << std::to_string(static_cast(coordinate_transform_mode)) << " " << cnode->fullname_with_scope(); return RET_ERROR; } @@ -187,7 +188,8 @@ STATUS ResizeMapper::Map(const CNodePtr &cnode, std::vector *ba if (prim->GetAttr(ops::kMethod) != nullptr) { auto interpolation_mode = static_cast(resize_prim->get_method()); if (kInterpolationModeMap.find(interpolation_mode) == kInterpolationModeMap.end()) { - MS_LOG(ERROR) << "unsupported interpolation mode:" << interpolation_mode << " " << cnode->fullname_with_scope(); + MS_LOG(ERROR) << "unsupported interpolation mode:" << std::to_string(static_cast(interpolation_mode)) << " " + << cnode->fullname_with_scope(); return RET_ERROR; } resize_operator->SetInterpolationMode(kInterpolationModeMap.at(interpolation_mode)); @@ -195,7 +197,8 @@ STATUS ResizeMapper::Map(const CNodePtr &cnode, std::vector *ba if (prim->GetAttr(ops::kNearestMode) != nullptr) { auto nearest_mode = static_cast(resize_prim->get_nearest_mode()); if (kNearestModeMap.find(nearest_mode) == kNearestModeMap.end()) { - MS_LOG(ERROR) << "unsupported nearest mode:" << nearest_mode << " " << cnode->fullname_with_scope(); + MS_LOG(ERROR) << "unsupported nearest mode:" << std::to_string(static_cast(nearest_mode)) << " " + << cnode->fullname_with_scope(); return RET_ERROR; } resize_operator->SetNearestMode(kNearestModeMap.at(nearest_mode)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.h index dea874e02c..a3a1039574 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/resize_mapper.h @@ -28,8 +28,8 @@ class ResizeMapper : public OpMapper { public: ResizeMapper() : OpMapper("Resize") {} ~ResizeMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.cc index 32423e1308..105286333c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.cc @@ -18,20 +18,19 @@ #include #include #include -#include "ops/op_utils.h" #include "common/anf_util.h" #include "ops/reverse_v2.h" #include "op/reverse_operator.h" namespace mindspore { namespace dpico { -STATUS ReverseMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ReverseMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto reverse_prim = utils::cast>(prim); + auto reverse_prim = api::utils::cast>(prim); MS_ASSERT(reverse_prim != nullptr); auto reverse_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.h index 4fc3d365c6..1be3e3010d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/reverse_mapper.h @@ -28,8 +28,8 @@ class ReverseMapper : public OpMapper { public: ReverseMapper() : OpMapper("Reverse") {} ~ReverseMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.cc index a25a80471b..76c8b43c13 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS RnnMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS RnnMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -42,25 +42,25 @@ STATUS RnnMapper::Map(const CNodePtr &cnode, std::vector *base_ } if (prim->GetAttr(kNumOutput) != nullptr) { - rnn_operator->SetRecurrentNumOutput(GetValue(prim->GetAttr(kNumOutput))); + rnn_operator->SetRecurrentNumOutput(static_cast(api::GetValue(prim->GetAttr(kNumOutput)))); } if (prim->GetAttr(kExposeHidden) != nullptr) { - rnn_operator->SetRecurrentExposeHidden(GetValue(prim->GetAttr(kExposeHidden))); + rnn_operator->SetRecurrentExposeHidden(api::GetValue(prim->GetAttr(kExposeHidden))); } if (prim->GetAttr(kOutputLastFrameFlag) != nullptr) { - rnn_operator->SetOutputLastFrameFlag(GetValue(prim->GetAttr(kOutputLastFrameFlag))); + rnn_operator->SetOutputLastFrameFlag(api::GetValue(prim->GetAttr(kOutputLastFrameFlag))); } if (prim->GetAttr(kInitialHOnlineFlag) != nullptr) { - rnn_operator->SetInitialHOnlineFlag(GetValue(prim->GetAttr(kInitialHOnlineFlag))); + rnn_operator->SetInitialHOnlineFlag(api::GetValue(prim->GetAttr(kInitialHOnlineFlag))); } if (prim->GetAttr(kUseDefaultInitialHFlag) != nullptr) { - rnn_operator->SetUseDefaultInitialHFlag(GetValue(prim->GetAttr(kUseDefaultInitialHFlag))); + rnn_operator->SetUseDefaultInitialHFlag(api::GetValue(prim->GetAttr(kUseDefaultInitialHFlag))); } if (prim->GetAttr(kKeepDirectionDimFlag) != nullptr) { - rnn_operator->SetKeepDirectionDimFlag(GetValue(prim->GetAttr(kKeepDirectionDimFlag))); + rnn_operator->SetKeepDirectionDimFlag(api::GetValue(prim->GetAttr(kKeepDirectionDimFlag))); } if (prim->GetAttr(kHasOutputGateFlag) != nullptr) { - rnn_operator->SetHasOutputGateFlag(GetValue(prim->GetAttr(kHasOutputGateFlag))); + rnn_operator->SetHasOutputGateFlag(api::GetValue(prim->GetAttr(kHasOutputGateFlag))); } if (SetRecurrentDataInfo(cnode, rnn_operator.get()) != RET_OK) { MS_LOG(ERROR) << "set rnn data info failed."; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.h index 45a5a2e7b7..ed15162130 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/rnn_mapper.h @@ -28,8 +28,8 @@ class RnnMapper : public OpMapper { public: RnnMapper() : OpMapper("Rnn") {} ~RnnMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.cc index 2d9a67b923..020b30de77 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.cc @@ -20,13 +20,12 @@ #include #include #include "common/op_attr.h" -#include "ops/op_utils.h" #include "op/roi_align_operator.h" namespace mindspore { namespace dpico { -STATUS RoiAlignMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS RoiAlignMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -45,7 +44,7 @@ STATUS RoiAlignMapper::Map(const CNodePtr &cnode, std::vector * roi_align_operator->SetOpType(mapper::OpType::ROIALIGN); if (prim->GetAttr(ops::kMode) != nullptr) { - auto pool_mode = GetValue(prim->GetAttr(ops::kMode)); + auto pool_mode = api::GetValue(prim->GetAttr(ops::kMode)); if (pool_mode == "avg") { roi_align_operator->SetPoolMode(mapper::RoiAlignPoolMode::ROI_ALIGN_AVG); } else if (pool_mode == "max") { @@ -56,16 +55,19 @@ STATUS RoiAlignMapper::Map(const CNodePtr &cnode, std::vector * } } if (prim->GetAttr(dpico::kOutputHeight) != nullptr) { - roi_align_operator->SetPooledHeight(GetValue(prim->GetAttr(dpico::kOutputHeight))); + roi_align_operator->SetPooledHeight( + static_cast(api::GetValue(prim->GetAttr(dpico::kOutputHeight)))); } if (prim->GetAttr(dpico::kOutputWidth) != nullptr) { - roi_align_operator->SetPooledWidth(GetValue(prim->GetAttr(dpico::kOutputWidth))); + roi_align_operator->SetPooledWidth( + static_cast(api::GetValue(prim->GetAttr(dpico::kOutputWidth)))); } if (prim->GetAttr(dpico::kSamplingRatio) != nullptr) { - roi_align_operator->SetSamplingRatio(GetValue(prim->GetAttr(dpico::kSamplingRatio))); + roi_align_operator->SetSamplingRatio( + static_cast(api::GetValue(prim->GetAttr(dpico::kSamplingRatio)))); } if (prim->GetAttr(dpico::kSpatialScale) != nullptr) { - roi_align_operator->SetSpatialScale(GetValue(prim->GetAttr(dpico::kSpatialScale))); + roi_align_operator->SetSpatialScale(api::GetValue(prim->GetAttr(dpico::kSpatialScale))); } base_operators->push_back(std::move(roi_align_operator)); return RET_OK; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.h index 73254c0502..a3766b867d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_align_mapper.h @@ -28,8 +28,8 @@ class RoiAlignMapper : public OpMapper { public: RoiAlignMapper() : OpMapper("RoiAlign") {} ~RoiAlignMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.cc index 3826a4fb8b..4a7abae9ff 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.cc @@ -18,19 +18,18 @@ #include #include #include -#include "ops/op_utils.h" #include "ops/roi_pooling.h" #include "op/roi_pool_operator.h" namespace mindspore { namespace dpico { -STATUS RoiPoolMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS RoiPoolMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto roi_pool_prim = utils::cast>(prim); + auto roi_pool_prim = api::utils::cast>(prim); MS_ASSERT(roi_pool_prim != nullptr); auto roi_pool_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.h index 240ae13d50..af13e4cf1d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/roi_pool_mapper.h @@ -28,8 +28,8 @@ class RoiPoolMapper : public OpMapper { public: RoiPoolMapper() : OpMapper("RoiPool") {} ~RoiPoolMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.cc index 824980ace0..c2bfd1d4a7 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.cc @@ -27,7 +27,7 @@ namespace mindspore { namespace dpico { namespace { -STATUS SetScaleDataInfo(const CNodePtr &cnode, mapper::ScaleOperator *scale_operator) { +STATUS SetScaleDataInfo(const api::CNodePtr &cnode, mapper::ScaleOperator *scale_operator) { if (scale_operator == nullptr) { MS_LOG(ERROR) << "scale_operator is nullptr."; return RET_ERROR; @@ -35,13 +35,13 @@ STATUS SetScaleDataInfo(const CNodePtr &cnode, mapper::ScaleOperator *scale_oper for (size_t i = 2; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); MS_ASSERT(input_node != nullptr); - auto param_node = input_node->cast(); + auto param_node = input_node->cast(); if (param_node == nullptr || !param_node->has_default()) { continue; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); MS_CHECK_TRUE_MSG(data != nullptr, RET_ERROR, "data is nullptr."); if (i == kInputIndex2) { scale_operator->SetScaleWeightPtr(data); @@ -69,13 +69,13 @@ STATUS SetScaleDataInfo(const CNodePtr &cnode, mapper::ScaleOperator *scale_oper return RET_OK; } } // namespace -STATUS ScaleMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ScaleMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto scale_prim = utils::cast>(prim); + auto scale_prim = api::utils::cast>(prim); MS_ASSERT(scale_prim != nullptr); auto scale_operator = std::make_unique(); @@ -92,10 +92,10 @@ STATUS ScaleMapper::Map(const CNodePtr &cnode, std::vector *bas scale_operator->SetOpType(mapper::OpType::SCALE); scale_operator->SetAxis(static_cast(scale_prim->get_axis())); if (scale_prim->GetAttr(kBiasTerm) != nullptr) { - scale_operator->SetScaleBiasFlag(GetValue(scale_prim->GetAttr(kBiasTerm))); + scale_operator->SetScaleBiasFlag(api::GetValue(scale_prim->GetAttr(kBiasTerm))); } if (scale_prim->GetAttr(kNumAxes) != nullptr) { - scale_operator->SetScaleNumAxes(GetValue(scale_prim->GetAttr(kNumAxes))); + scale_operator->SetScaleNumAxes(static_cast(api::GetValue(scale_prim->GetAttr(kNumAxes)))); } if (SetScaleDataInfo(cnode, scale_operator.get()) != RET_OK) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.h index ba8ef52c38..2a2e256a8d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/scale_mapper.h @@ -28,8 +28,8 @@ class ScaleMapper : public OpMapper { public: ScaleMapper() : OpMapper("Scale") {} ~ScaleMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.cc index 11f77fbe35..ec9f86a841 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS ShapeMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS ShapeMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.h index 59d18ac13c..dc21d92ce1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/shape_mapper.h @@ -28,8 +28,8 @@ class ShapeMapper : public OpMapper { public: ShapeMapper() : OpMapper("Shape") {} ~ShapeMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.cc index 2cbe3837db..67eef081d2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.cc @@ -18,14 +18,13 @@ #include #include #include -#include "ops/op_utils.h" #include "common/anf_util.h" #include "op/shuffle_channel_operator.h" namespace mindspore { namespace dpico { -STATUS ShuffleChannelMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS ShuffleChannelMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -44,7 +43,8 @@ STATUS ShuffleChannelMapper::Map(const CNodePtr &cnode, std::vectorSetOpType(mapper::OpType::SHUFFLECHANNEL); if (prim->GetAttr(ops::kGroup) != nullptr) { - shuffle_channel_operator->SetShuffleChannelGroup(GetValue(prim->GetAttr(ops::kGroup))); + shuffle_channel_operator->SetShuffleChannelGroup( + static_cast(api::GetValue(prim->GetAttr(ops::kGroup)))); } base_operators->push_back(std::move(shuffle_channel_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.h index 7e036dfe6d..2520cd4b81 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/shuffle_channel_mapper.h @@ -28,8 +28,8 @@ class ShuffleChannelMapper : public OpMapper { public: ShuffleChannelMapper() : OpMapper("ShuffleChannel") {} ~ShuffleChannelMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.cc index 529607ec5f..13de9c7024 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.cc @@ -21,19 +21,18 @@ #include #include "common/anf_util.h" #include "common/op_enum.h" -#include "ops/op_utils.h" #include "ops/split.h" #include "op/slice_operator.h" namespace mindspore { namespace dpico { -STATUS SliceMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS SliceMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto split_prim = utils::cast>(prim); + auto split_prim = api::utils::cast>(prim); MS_ASSERT(split_prim != nullptr); auto slice_operator = std::make_unique(); @@ -80,10 +79,12 @@ STATUS SliceMapper::Map(const CNodePtr &cnode, std::vector *bas MS_LOG(ERROR) << "split sizes is invalid, which is not larger than 0 and less than uint32_max"; return RET_ERROR; } - std::vector sizes_u(sizes.begin(), sizes.end()); + std::vector sizes_u; + (void)std::transform(sizes.begin(), sizes.end(), std::back_inserter(sizes_u), + [](int64_t size) { return static_cast(size); }); uint32_t slice_point_cnt = 0; for (size_t i = 0; i < sizes_u.size() - 1; i++) { - if (sizes_u.at(i) >= static_cast(shape[split_axis]) - slice_point_cnt) { + if (sizes_u.at(i) >= (static_cast(shape[split_axis]) - slice_point_cnt)) { MS_LOG(ERROR) << "split sizes is invalid, which is larger than the related dim."; return RET_ERROR; } @@ -97,7 +98,7 @@ STATUS SliceMapper::Map(const CNodePtr &cnode, std::vector *bas MS_LOG(ERROR) << "cannot determine split points."; return RET_ERROR; } - auto output_num = GetValue(split_prim->GetAttr(ops::kOutputNum)); + auto output_num = api::GetValue(split_prim->GetAttr(ops::kOutputNum)); if (shape[split_axis] % output_num != 0) { MS_LOG(ERROR) << "split op is invalid, which input shape cannot be splited."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.h index dbcc1d4d9a..75a8ab1692 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/slice_mapper.h @@ -28,8 +28,8 @@ class SliceMapper : public OpMapper { public: SliceMapper() : OpMapper("Slice") {} ~SliceMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.cc index c1d7ec9cd8..b8d9b6b01a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.cc @@ -25,13 +25,13 @@ namespace mindspore { namespace dpico { -STATUS SoftmaxMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS SoftmaxMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto softmax_prim = utils::cast>(prim); + auto softmax_prim = api::utils::cast>(prim); MS_ASSERT(softmax_prim != nullptr); auto softmax_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.h index ac7e60bd22..c4772ffb31 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/softmax_mapper.h @@ -28,8 +28,8 @@ class SoftmaxMapper : public OpMapper { public: SoftmaxMapper() : OpMapper("Softmax") {} ~SoftmaxMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.cc index 50dbe48dfb..f4781c9b88 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.cc @@ -17,17 +17,15 @@ #include "mapper/spp_mapper.h" #include #include -#include #include #include "common/op_attr.h" #include "common/anf_util.h" -#include "ops/op_utils.h" #include "op/spp_operator.h" namespace mindspore { namespace dpico { -STATUS SppMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS SppMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -47,10 +45,10 @@ STATUS SppMapper::Map(const CNodePtr &cnode, std::vector *base_ spp_operator->SetOpType(mapper::OpType::SPPPOOLING); if (prim->GetAttr(dpico::kPyramidHeight) != nullptr) { - spp_operator->SetPyramidHeight(static_cast(GetValue(prim->GetAttr(dpico::kPyramidHeight)))); + spp_operator->SetPyramidHeight(static_cast(api::GetValue(prim->GetAttr(dpico::kPyramidHeight)))); } if (prim->GetAttr(dpico::kPoolMethod) != nullptr) { - auto pool_method = GetValue(prim->GetAttr(dpico::kPoolMethod)); + auto pool_method = api::GetValue(prim->GetAttr(dpico::kPoolMethod)); if (pool_method == 0) { spp_operator->SetSppType(mapper::OpType::POOLINGMAX); } else if (pool_method == 1) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.h index 3f73a65d7f..5a2fe84d07 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/spp_mapper.h @@ -28,8 +28,8 @@ class SppMapper : public OpMapper { public: SppMapper() : OpMapper("Spp") {} ~SppMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.cc index 93fb603ef6..ccd775b0cb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.cc @@ -19,20 +19,18 @@ #include #include #include -#include "common/anf_util.h" -#include "ops/op_utils.h" #include "ops/squeeze.h" #include "op/squeeze_operator.h" namespace mindspore { namespace dpico { -STATUS SqueezeMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS SqueezeMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto squeeze_prim = utils::cast>(prim); + auto squeeze_prim = api::utils::cast>(prim); MS_ASSERT(squeeze_prim != nullptr); auto squeeze_operator = std::make_unique(); @@ -49,7 +47,7 @@ STATUS SqueezeMapper::Map(const CNodePtr &cnode, std::vector *b squeeze_operator->SetOpType(mapper::OpType::SQUEEZE); if (squeeze_prim->GetAttr(ops::kAxis) != nullptr) { - auto axes = GetValue>(squeeze_prim->GetAttr(ops::kAxis)); + auto axes = api::GetValue>(squeeze_prim->GetAttr(ops::kAxis)); std::vector dims; (void)std::transform(axes.begin(), axes.end(), std::back_inserter(dims), [](const int64_t &value) { return static_cast(value); }); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.h index 5cbb2ee04a..ad5cf088c2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/squeeze_mapper.h @@ -28,8 +28,8 @@ class SqueezeMapper : public OpMapper { public: SqueezeMapper() : OpMapper("Squeeze") {} ~SqueezeMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.cc index 18e64b856b..eb722ee63a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.cc @@ -20,14 +20,13 @@ #include #include "common/anf_util.h" #include "common/op_enum.h" -#include "ops/op_utils.h" #include "ops/strided_slice.h" #include "op/extract_slice_operator.h" namespace mindspore { namespace dpico { namespace { -STATUS SetExtractSliceDataInfo(const CNodePtr &cnode, mapper::ExtractSliceOperator *extract_slice_operator) { +STATUS SetExtractSliceDataInfo(const api::CNodePtr &cnode, mapper::ExtractSliceOperator *extract_slice_operator) { if (extract_slice_operator == nullptr) { MS_LOG(ERROR) << "extract_slice_operator is nullptr."; return RET_ERROR; @@ -35,13 +34,13 @@ STATUS SetExtractSliceDataInfo(const CNodePtr &cnode, mapper::ExtractSliceOperat for (size_t i = 2; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); MS_ASSERT(input_node != nullptr); - auto param_node = input_node->cast(); + auto param_node = input_node->cast(); if (param_node == nullptr || !param_node->has_default()) { continue; } - auto tensor_info = std::dynamic_pointer_cast(param_node->default_param()); + auto tensor_info = param_node->default_param()->cast(); if (tensor_info != nullptr && tensor_info->DataSize() != 0) { - auto data = reinterpret_cast(tensor_info->data_c()); + auto data = reinterpret_cast(tensor_info->data()); if (i == kInputIndex2) { extract_slice_operator->SetStartsVec(std::vector(data, data + tensor_info->DataSize())); } else if (i == kInputIndex3) { @@ -64,13 +63,13 @@ STATUS SetExtractSliceDataInfo(const CNodePtr &cnode, mapper::ExtractSliceOperat return RET_OK; } } // namespace -STATUS StridedSliceMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS StridedSliceMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto strided_slice_prim = utils::cast>(prim); + auto strided_slice_prim = api::utils::cast>(prim); MS_ASSERT(strided_slice_prim != nullptr); auto extract_slice_operator = std::make_unique(); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.h index 78c2da6a39..881e2cde2d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/strided_slice_mapper.h @@ -28,8 +28,8 @@ class StridedSliceMapper : public OpMapper { public: StridedSliceMapper() : OpMapper("StridedSlice") {} ~StridedSliceMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.cc index f874cb8668..0ec68c057b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.cc @@ -24,8 +24,8 @@ namespace mindspore { namespace dpico { -STATUS ThresholdMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS ThresholdMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -44,7 +44,7 @@ STATUS ThresholdMapper::Map(const CNodePtr &cnode, std::vector threshold_operator->SetOpType(mapper::OpType::THRESHOLD); if (prim->GetAttr(dpico::kThreshold) != nullptr) { - threshold_operator->SetThreshold(GetValue(prim->GetAttr(dpico::kThreshold))); + threshold_operator->SetThreshold(api::GetValue(prim->GetAttr(dpico::kThreshold))); } base_operators->push_back(std::move(threshold_operator)); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.h index bf6d214b78..2bf7f8b070 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/threshold_mapper.h @@ -28,8 +28,8 @@ class ThresholdMapper : public OpMapper { public: ThresholdMapper() : OpMapper("Threshold") {} ~ThresholdMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.cc index 945289f88e..58dfb966f8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.cc @@ -28,13 +28,13 @@ namespace mindspore { namespace dpico { -STATUS TileMapper::Map(const CNodePtr &cnode, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) { +STATUS TileMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto tile_prim = utils::cast>(prim); + auto tile_prim = api::utils::cast>(prim); MS_ASSERT(tile_prim != nullptr); auto tile_operator = std::make_unique(); @@ -50,7 +50,7 @@ STATUS TileMapper::Map(const CNodePtr &cnode, std::vector *base tile_operator->SetOpType(mapper::OpType::TILE); if (tile_prim->GetAttr(ops::kDims) != nullptr) { - auto dims = GetValue>(tile_prim->GetAttr(ops::kDims)); + auto dims = api::GetValue>(tile_prim->GetAttr(ops::kDims)); if (dims.size() == 1) { // tf tile has multiple axis tile_operator->SetAxis(dims.at(0)); } @@ -78,9 +78,7 @@ STATUS TileMapper::Map(const CNodePtr &cnode, std::vector *base [](const int64_t &value) { return static_cast(value); }); } else if (tile_prim->GetAttr(dpico::kMultiples) != nullptr) { if (tile_prim->GetAttr(ops::kDims) != nullptr) { - auto data = GetValue>(tile_prim->GetAttr(dpico::kMultiples)); - (void)std::transform(data.begin(), data.end(), std::back_inserter(tiles), - [](const int64_t &value) { return static_cast(value); }); + tiles = api::GetValue>(tile_prim->GetAttr(dpico::kMultiples)); } } else { MS_LOG(ERROR) << "null param"; diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.h index 529bfcc1b2..4ee57bacf9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/tile_mapper.h @@ -28,8 +28,8 @@ class TileMapper : public OpMapper { public: TileMapper() : OpMapper("Tile") {} ~TileMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.cc index fb9fda32f4..a6c522c17e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.cc @@ -20,19 +20,18 @@ #include #include #include "common/anf_util.h" -#include "ops/op_utils.h" #include "ops/unsqueeze.h" #include "op/unsqueeze_operator.h" namespace mindspore { namespace dpico { -STATUS UnsqueezeMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS UnsqueezeMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; } - auto unsqueeze_prim = utils::cast>(prim); + auto unsqueeze_prim = api::utils::cast>(prim); MS_ASSERT(unsqueeze_prim != nullptr); auto unsqueeze_operator = std::make_unique(); @@ -49,7 +48,7 @@ STATUS UnsqueezeMapper::Map(const CNodePtr &cnode, std::vector unsqueeze_operator->SetOpType(mapper::OpType::UNSQUEEZE); if (unsqueeze_prim->GetAttr(ops::kAxis) != nullptr) { - auto axes = GetValue>(unsqueeze_prim->GetAttr(ops::kAxis)); + auto axes = api::GetValue>(unsqueeze_prim->GetAttr(ops::kAxis)); std::vector dims; (void)std::transform(axes.begin(), axes.end(), std::back_inserter(dims), [](const int64_t &value) { return static_cast(value); }); diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.h index af3972bb5b..e0e60225b9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/unsqueeze_mapper.h @@ -28,8 +28,8 @@ class UnsqueezeMapper : public OpMapper { public: UnsqueezeMapper() : OpMapper("Unsqueeze") {} ~UnsqueezeMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.cc b/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.cc index c4fd089ccf..9086345b0f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.cc @@ -20,13 +20,12 @@ #include #include #include "common/op_attr.h" -#include "ops/op_utils.h" #include "op/upsample_operator.h" namespace mindspore { namespace dpico { -STATUS UpsampleMapper::Map(const CNodePtr &cnode, std::vector *base_operators, - const PrimitivePtr &prim, const CNodePtrList &output_cnodes) { +STATUS UpsampleMapper::Map(const api::CNodePtr &cnode, std::vector *base_operators, + const api::PrimitivePtr &prim, const api::CNodePtrList &output_cnodes) { if (base_operators == nullptr) { MS_LOG(ERROR) << "base_operators is nullptr."; return RET_ERROR; @@ -45,16 +44,16 @@ STATUS UpsampleMapper::Map(const CNodePtr &cnode, std::vector * upsample_operator->SetOpType(mapper::OpType::UPSAMPLE); if (prim->GetAttr(ops::kScale) != nullptr) { - upsample_operator->SetUpsampleScale(GetValue(prim->GetAttr(ops::kScale))); + upsample_operator->SetUpsampleScale(api::GetValue(prim->GetAttr(ops::kScale))); } if (prim->GetAttr(kUpsampleH) != nullptr) { - upsample_operator->SetUpsampleHeight(GetValue(prim->GetAttr(kUpsampleH))); + upsample_operator->SetUpsampleHeight(static_cast(api::GetValue(prim->GetAttr(kUpsampleH)))); } if (prim->GetAttr(kUpsampleW) != nullptr) { - upsample_operator->SetUpsampleWidth(GetValue(prim->GetAttr(kUpsampleW))); + upsample_operator->SetUpsampleWidth(static_cast(api::GetValue(prim->GetAttr(kUpsampleW)))); } if (prim->GetAttr(kInterpolationMode) != nullptr) { - auto mode = GetValue(prim->GetAttr(kInterpolationMode)); + auto mode = api::GetValue(prim->GetAttr(kInterpolationMode)); if (mode == kNearest) { upsample_operator->SetInterpolationMode(mapper::InterpolationMode::NEAREST); } else if (mode == kBilinear) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.h b/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.h index a890a68628..3c0a0c5086 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.h +++ b/mindspore/lite/tools/converter/adapter/dpico/mapper/upsample_mapper.h @@ -28,8 +28,8 @@ class UpsampleMapper : public OpMapper { public: UpsampleMapper() : OpMapper("Upsample") {} ~UpsampleMapper() override = default; - STATUS Map(const CNodePtr &node, std::vector *base_operators, const PrimitivePtr &prim, - const CNodePtrList &output_cnodes) override; + STATUS Map(const api::CNodePtr &node, std::vector *base_operators, const api::PrimitivePtr &prim, + const api::CNodePtrList &output_cnodes) override; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.cc index 338b5e7d2e..bf79a75e5f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.cc @@ -17,18 +17,17 @@ #include "parser/caffe/caffe_absval_parser.h" #include #include -#include "common/op_attr.h" #include "ops/abs.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeAbsvalParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeAbsvalParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeAbsvalParser("AbsVal", new CaffeAbsvalParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.h index f9a44b2567..9008997876 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_absval_parser.h @@ -28,7 +28,7 @@ class CaffeAbsvalParser : public CaffeNodeParser { CaffeAbsvalParser() : CaffeNodeParser("Absval") {} ~CaffeAbsvalParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.cc index e1c85307b2..f13e98f962 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.cc @@ -17,13 +17,14 @@ #include "parser/caffe/caffe_activation_parser.h" #include #include +#include #include "common/op_attr.h" #include "ops/fusion/activation.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -38,11 +39,11 @@ ops::PrimitiveC *CaffeReluParser::Parse(const caffe::LayerParameter &proto, cons } } - return prim.release(); + return prim; } -ops::PrimitiveC *CaffeRelu6Parser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeRelu6Parser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -54,33 +55,33 @@ ops::PrimitiveC *CaffeRelu6Parser::Parse(const caffe::LayerParameter &proto, con prim->set_alpha(negative_slope); } } - return prim.release(); + return prim; } -ops::PrimitiveC *CaffeSigmoidParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeSigmoidParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; } prim->set_activation_type(mindspore::ActivationType::SIGMOID); - return prim.release(); + return prim; } -ops::PrimitiveC *CaffeTanhParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeTanhParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; } prim->set_activation_type(mindspore::ActivationType::TANH); - return prim.release(); + return prim; } -ops::PrimitiveC *CaffeEluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeEluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -94,11 +95,11 @@ ops::PrimitiveC *CaffeEluParser::Parse(const caffe::LayerParameter &proto, const } } - return prim.release(); + return prim; } -ops::PrimitiveC *CaffeHswishParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeHswishParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -107,10 +108,10 @@ ops::PrimitiveC *CaffeHswishParser::Parse(const caffe::LayerParameter &proto, co if (proto.has_hswish_param()) { const caffe::HswishParameter &hswishParameter = proto.hswish_param(); if (hswishParameter.has_negative_slope()) { - prim->AddAttr(dpico::kNegativeSlope, MakeValue(hswishParameter.negative_slope())); + prim->AddAttr(dpico::kNegativeSlope, api::MakeValue(hswishParameter.negative_slope())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeReluParser("ReLU", new CaffeReluParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.h index bfe9cb1658..5e40e26579 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_activation_parser.h @@ -28,7 +28,7 @@ class CaffeReluParser : public CaffeNodeParser { CaffeReluParser() : CaffeNodeParser("relu") {} ~CaffeReluParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeRelu6Parser : public CaffeNodeParser { @@ -36,7 +36,7 @@ class CaffeRelu6Parser : public CaffeNodeParser { CaffeRelu6Parser() : CaffeNodeParser("relu6") {} ~CaffeRelu6Parser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeSigmoidParser : public CaffeNodeParser { @@ -44,7 +44,7 @@ class CaffeSigmoidParser : public CaffeNodeParser { CaffeSigmoidParser() : CaffeNodeParser("sigmoid") {} ~CaffeSigmoidParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeTanhParser : public CaffeNodeParser { @@ -52,7 +52,7 @@ class CaffeTanhParser : public CaffeNodeParser { CaffeTanhParser() : CaffeNodeParser("tanh") {} ~CaffeTanhParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeEluParser : public CaffeNodeParser { @@ -60,7 +60,7 @@ class CaffeEluParser : public CaffeNodeParser { CaffeEluParser() : CaffeNodeParser("elu") {} ~CaffeEluParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeHswishParser : public CaffeNodeParser { @@ -68,7 +68,7 @@ class CaffeHswishParser : public CaffeNodeParser { CaffeHswishParser() : CaffeNodeParser("Hswish") {} ~CaffeHswishParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.cc index 7c9fa248a5..0f8d71ddca 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.cc @@ -20,8 +20,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeArgMaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeArgMaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -41,7 +41,7 @@ ops::PrimitiveC *CaffeArgMaxParser::Parse(const caffe::LayerParameter &proto, co prim->set_axis(argmaxParam.axis()); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeArgMaxParser("ArgMax", new CaffeArgMaxParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.h index 8bbfc0d19c..d1a970e060 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_argmax_parser.h @@ -28,7 +28,7 @@ class CaffeArgMaxParser : public CaffeNodeParser { CaffeArgMaxParser() : CaffeNodeParser("argmax") {} ~CaffeArgMaxParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.cc index 07250c6030..2c3298ba2b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.cc @@ -22,8 +22,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeBatchNormParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeBatchNormParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -50,10 +50,10 @@ ops::PrimitiveC *CaffeBatchNormParser::Parse(const caffe::LayerParameter &proto, prim->set_epsilon(epsilon); if (batchNormParam.has_use_global_stats()) { - prim->AddAttr(dpico::kUseGlobalStats, MakeValue(batchNormParam.use_global_stats())); + prim->AddAttr(dpico::kUseGlobalStats, api::MakeValue(batchNormParam.use_global_stats())); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeBatchNormParser("BatchNorm", new CaffeBatchNormParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.h index 4ddc7b260b..ae0ce14b1a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_batchnorm_parser.h @@ -28,7 +28,7 @@ class CaffeBatchNormParser : public CaffeNodeParser { CaffeBatchNormParser() : CaffeNodeParser("batchnorm") {} ~CaffeBatchNormParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.cc index 85ebeafd28..cf0fc52b16 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.cc @@ -23,8 +23,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeBiLstmParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeBiLstmParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -33,15 +33,15 @@ ops::PrimitiveC *CaffeBiLstmParser::Parse(const caffe::LayerParameter &proto, co if (proto.has_recurrent_param()) { const auto &bi_lstm_param = proto.recurrent_param(); if (bi_lstm_param.has_num_output()) { - prim->AddAttr(dpico::kNumOutput, MakeValue(bi_lstm_param.num_output())); - prim->AddAttr(dpico::kOutputChannel, MakeValue(bi_lstm_param.num_output())); + prim->AddAttr(dpico::kNumOutput, api::MakeValue(bi_lstm_param.num_output())); + prim->AddAttr(dpico::kOutputChannel, api::MakeValue(bi_lstm_param.num_output())); } if (bi_lstm_param.has_expose_hidden()) { - prim->AddAttr(dpico::kExposeHidden, MakeValue(bi_lstm_param.expose_hidden())); + prim->AddAttr(dpico::kExposeHidden, api::MakeValue(bi_lstm_param.expose_hidden())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeBiLstmParser("BILSTM", new CaffeBiLstmParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.h index e138dff94f..670599468f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bi_lstm_parser.h @@ -28,7 +28,7 @@ class CaffeBiLstmParser : public CaffeNodeParser { CaffeBiLstmParser() : CaffeNodeParser("BiLstm") {} ~CaffeBiLstmParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.cc index 64ea9390ba..c31ea359d3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.cc @@ -19,11 +19,12 @@ #include #include "common/op_attr.h" #include "ops/custom.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeBiasParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeBiasParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -31,13 +32,13 @@ ops::PrimitiveC *CaffeBiasParser::Parse(const caffe::LayerParameter &proto, cons prim->set_type("Bias"); const caffe::BiasParameter &biasParam = proto.bias_param(); if (biasParam.has_axis()) { - prim->AddAttr(ops::kAxis, MakeValue(biasParam.axis())); + prim->AddAttr(ops::kAxis, api::MakeValue(biasParam.axis())); } if (biasParam.has_num_axes()) { - prim->AddAttr(dpico::kNumAxes, MakeValue(biasParam.num_axes())); + prim->AddAttr(dpico::kNumAxes, api::MakeValue(biasParam.num_axes())); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeBiasParser("Bias", new CaffeBiasParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.h index 8118ff0905..e676e1cc12 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bias_parser.h @@ -28,7 +28,7 @@ class CaffeBiasParser : public CaffeNodeParser { CaffeBiasParser() : CaffeNodeParser("Bias") {} ~CaffeBiasParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.cc index 0a1bd119e8..dc03ea3ed3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.cc @@ -28,35 +28,35 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeBinaryMathParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +BaseOperatorPtr CaffeBinaryMathParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { const caffe::BinaryMathParameter &binary_math_param = proto.binary_math_param(); - std::unique_ptr prim; + BaseOperatorPtr prim; if (binary_math_param.has_operation()) { auto operation = binary_math_param.operation(); switch (operation) { case caffe::BinaryMathParameter_BinaryMathOp_ADD: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_SUB: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_MUL: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_DIV: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_MAX: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_MIN: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_SQUARE_DIFF: - prim = std::make_unique(); + prim = std::make_shared(); break; case caffe::BinaryMathParameter_BinaryMathOp_X_DIV_Y: { - prim = std::make_unique(); + prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr. " << proto.name(); return nullptr; @@ -65,7 +65,7 @@ ops::PrimitiveC *CaffeBinaryMathParser::Parse(const caffe::LayerParameter &proto break; } case caffe::BinaryMathParameter_BinaryMathOp_X_LOG_Y: { - prim = std::make_unique(); + prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr. " << proto.name(); return nullptr; @@ -78,13 +78,13 @@ ops::PrimitiveC *CaffeBinaryMathParser::Parse(const caffe::LayerParameter &proto return nullptr; } } else { - prim = std::make_unique(); + prim = std::make_shared(); } if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr. " << proto.name(); return nullptr; } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeBinaryMathParser("BinaryMath", new CaffeBinaryMathParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.h index edcf73efad..3077ac826c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_binary_math_parser.h @@ -28,7 +28,7 @@ class CaffeBinaryMathParser : public CaffeNodeParser { CaffeBinaryMathParser() : CaffeNodeParser("binary_math") {} ~CaffeBinaryMathParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.cc index fce5ad5ce8..bb290237cd 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.cc @@ -22,15 +22,15 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeBnllParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeBnllParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; } prim->set_type("Bnll"); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeBnllParser("BNLL", new CaffeBnllParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.h index d8387397d7..b2a50304a3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_bnll_parser.h @@ -28,7 +28,7 @@ class CaffeBnllParser : public CaffeNodeParser { CaffeBnllParser() : CaffeNodeParser("Bnll") {} ~CaffeBnllParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.cc index 2204b3df01..b3dae2cfc1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.cc @@ -21,8 +21,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeClipParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeClipParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -34,7 +34,7 @@ ops::PrimitiveC *CaffeClipParser::Parse(const caffe::LayerParameter &proto, cons if (clipParam.has_max()) { prim->set_max(clipParam.max()); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeClipParser("Clip", new CaffeClipParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.h index 1de7808254..6cfcc37e3f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_clip_parser.h @@ -28,7 +28,7 @@ class CaffeClipParser : public CaffeNodeParser { CaffeClipParser() : CaffeNodeParser("Clip") {} ~CaffeClipParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.cc index 99e4d7735d..8a2dae1cf9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.cc @@ -20,8 +20,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeConcatParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeConcatParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -46,7 +46,7 @@ ops::PrimitiveC *CaffeConcatParser::Parse(const caffe::LayerParameter &proto, co } prim->set_axis(axis); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeConcatParser("Concat", new CaffeConcatParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.h index 10dbd4e84b..3f94a18c75 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_concat_parser.h @@ -28,7 +28,7 @@ class CaffeConcatParser : public CaffeNodeParser { CaffeConcatParser() : CaffeNodeParser("concat") {} ~CaffeConcatParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.cc index 5df87e306f..ae3c8d8b30 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.cc @@ -22,9 +22,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeConvolutionParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeConvolutionParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -106,10 +105,10 @@ ops::PrimitiveC *CaffeConvolutionParser::Parse(const caffe::LayerParameter &prot } if (convParam.has_bias_term()) { - prim->AddAttr(dpico::kBiasTerm, MakeValue(convParam.bias_term())); + prim->AddAttr(dpico::kBiasTerm, api::MakeValue(convParam.bias_term())); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeConvolutionParser("Convolution", new CaffeConvolutionParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.h index 1cac984f87..12d829e551 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_convolution_parser.h @@ -30,7 +30,7 @@ class CaffeConvolutionParser : public CaffeNodeParser { explicit CaffeConvolutionParser(std::string nodeName) : CaffeNodeParser(std::move(nodeName)) {} ~CaffeConvolutionParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.cc index 8976f149cd..dbd8a22820 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.cc @@ -21,8 +21,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeCropParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeCropParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -52,7 +52,7 @@ ops::PrimitiveC *CaffeCropParser::Parse(const caffe::LayerParameter &proto, cons } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeCropParser("Crop", new CaffeCropParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.h index 818557f502..adc1c2b8a2 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_crop_parser.h @@ -28,7 +28,7 @@ class CaffeCropParser : public CaffeNodeParser { CaffeCropParser() : CaffeNodeParser("crop") {} ~CaffeCropParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.cc index 76a3b8780c..3858e12592 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.cc @@ -20,14 +20,14 @@ #include #include #include "common/op_attr.h" -#include "parser/caffe/caffe_detection_output_parser.h" -#include "parser/detection_output_param_holder.h" #include "ops/custom.h" +#include "third_party/securec/include/securec.h" +#include "parser/detection_output_param_helper.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeDecBBoxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeDecBBoxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -35,17 +35,17 @@ ops::PrimitiveC *CaffeDecBBoxParser::Parse(const caffe::LayerParameter &proto, c prim->set_type("DecBBox"); if (proto.has_num_anchors()) { - prim->AddAttr(dpico::kNumAnchors, MakeValue(proto.num_anchors())); + prim->AddAttr(dpico::kNumAnchors, api::MakeValue(proto.num_anchors())); } if (proto.has_num_bboxes_per_grid()) { - prim->AddAttr(dpico::kNumBboxesPerGrid, MakeValue(proto.num_bboxes_per_grid())); + prim->AddAttr(dpico::kNumBboxesPerGrid, api::MakeValue(proto.num_bboxes_per_grid())); } if (proto.has_num_coords()) { - prim->AddAttr(dpico::kNumCoords, MakeValue(proto.num_coords())); + prim->AddAttr(dpico::kNumCoords, api::MakeValue(proto.num_coords())); } if (proto.has_num_classes()) { - auto num_classes = proto.num_classes(); - prim->AddAttr(dpico::kNumClasses, MakeValue(num_classes)); + uint32_t num_classes = proto.num_classes(); + prim->AddAttr(dpico::kNumClasses, api::MakeValue(num_classes)); std::map> custom_attrs; std::vector num_classes_attr(sizeof(uint32_t)); if (memcpy_s(num_classes_attr.data(), num_classes_attr.size() * sizeof(uint8_t), &num_classes, sizeof(uint32_t)) != @@ -57,31 +57,17 @@ ops::PrimitiveC *CaffeDecBBoxParser::Parse(const caffe::LayerParameter &proto, c prim->set_attr(custom_attrs); } if (proto.has_num_grids_height()) { - prim->AddAttr(dpico::kNumGridsHeight, MakeValue(proto.num_grids_height())); + prim->AddAttr(dpico::kNumGridsHeight, api::MakeValue(proto.num_grids_height())); } if (proto.has_num_grids_width()) { - prim->AddAttr(dpico::kNumGridsWidth, MakeValue(proto.num_grids_width())); + prim->AddAttr(dpico::kNumGridsWidth, api::MakeValue(proto.num_grids_width())); } - const auto &decbbox_param = proto.decbbox_param(); - mapper::DetectionOutputParam param; - if (SetParamType(¶m, decbbox_param.param_type()) != RET_OK) { - MS_LOG(ERROR) << "Set param type failed."; + if (dpico::SetAttrsByDecBboxParam(prim, proto) != RET_OK) { + MS_LOG(ERROR) << "set attrs by dec bbox param"; return nullptr; } - if (SetCodeType(¶m, decbbox_param.code_type()) != RET_OK) { - MS_LOG(ERROR) << "Set code type failed."; - return nullptr; - } - (void)SetParamAttr(¶m, decbbox_param); - auto param_holder_ptr = std::make_shared(param); - if (param_holder_ptr == nullptr) { - MS_LOG(ERROR) << "new DetectionOutputParamHolder failed."; - return nullptr; - } - - prim->AddAttr(dpico::kDecBBoxParam, MakeValue(param_holder_ptr)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeDecBBoxParser("DecBBox", new CaffeDecBBoxParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.h index b4a5f40d8c..2a088f742a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_decbbox_parser.h @@ -28,7 +28,7 @@ class CaffeDecBBoxParser : public CaffeNodeParser { CaffeDecBBoxParser() : CaffeNodeParser("decbbox") {} ~CaffeDecBBoxParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.cc index 311bc46d69..19c216533f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.cc @@ -19,12 +19,13 @@ #include "common/op_enum.h" #include "common/check_base.h" #include "ops/fusion/conv2d_transpose_fusion.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeDeconvolutionParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeDeconvolutionParser::Parse(const caffe::LayerParameter &proto, + const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -104,10 +105,10 @@ ops::PrimitiveC *CaffeDeconvolutionParser::Parse(const caffe::LayerParameter &pr } } if (group != 1) { - prim->AddAttr(ops::kIsDepthWise, MakeValue(true)); + prim->AddAttr(ops::kIsDepthWise, api::MakeValue(true)); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeDeconvolutionParser("Deconvolution", new CaffeDeconvolutionParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.h index ce7473a56e..aa87de9117 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_deconvolution_parser.h @@ -27,7 +27,7 @@ class CaffeDeconvolutionParser : public CaffeNodeParser { CaffeDeconvolutionParser() : CaffeNodeParser("deconvolution") {} ~CaffeDeconvolutionParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.cc index 8829b630b9..ce04faecfe 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.cc @@ -17,14 +17,15 @@ #include "parser/caffe/caffe_depthwise_conv_parser.h" #include #include "ops/fusion/conv2d_fusion.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeDepthwiseConvolutionParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { +BaseOperatorPtr CaffeDepthwiseConvolutionParser::Parse(const caffe::LayerParameter &proto, + const caffe::LayerParameter &weight) { auto prim = CaffeConvolutionParser::Parse(proto, weight); if (prim != nullptr) { - prim->AddAttr(ops::kIsDepthWise, MakeValue(true)); + prim->AddAttr(ops::kIsDepthWise, api::MakeValue(true)); } return prim; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.h index c1f8086716..369516bcd9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_depthwise_conv_parser.h @@ -27,7 +27,7 @@ class CaffeDepthwiseConvolutionParser : public CaffeConvolutionParser { CaffeDepthwiseConvolutionParser() : CaffeConvolutionParser("depthwise_convolution") {} ~CaffeDepthwiseConvolutionParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.cc index a7a0efbce4..11fd28c118 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.cc @@ -18,87 +18,14 @@ #include #include #include "common/op_attr.h" -#include "parser/detection_output_param_holder.h" #include "ops/custom.h" +#include "parser/detection_output_param_helper.h" namespace mindspore { namespace lite { -STATUS SetParamType(mapper::DetectionOutputParam *param, const caffe::DetectionOutputParameter_ParamType ¶m_type) { - if (param == nullptr) { - MS_LOG(ERROR) << "input param is nullptr. "; - return RET_ERROR; - } - switch (param_type) { - case caffe::DetectionOutputParameter_ParamType_DecBBox: - param->paramType = mapper::ProposalParamType::PROPOSAL_DECBBOX; - break; - case caffe::DetectionOutputParameter_ParamType_Sort: - param->paramType = mapper::ProposalParamType::PROPOSAL_SORT; - break; - case caffe::DetectionOutputParameter_ParamType_Nms: - param->paramType = mapper::ProposalParamType::PROPOSAL_NMS; - break; - case caffe::DetectionOutputParameter_ParamType_FilterBox: - param->paramType = mapper::ProposalParamType::PROPOSAL_FILTERBOX; - break; - default: - MS_LOG(ERROR) << "Unsupported Param Type: " << param_type; - return RET_ERROR; - } - return RET_OK; -} -STATUS SetCodeType(mapper::DetectionOutputParam *param, const caffe::CodeType &code_type) { - if (param == nullptr) { - MS_LOG(ERROR) << "input param is nullptr. "; - return RET_ERROR; - } - if (code_type == caffe::CodeType::CENTER_SIZE) { - param->codeType = mapper::DecBboxCodeType::DECBBOX_CODE_TYPE_CENTER_SIZE; - } else { - MS_LOG(ERROR) << "Unsupported Code type: " << code_type; - return RET_ERROR; - } - return RET_OK; -} -void SetParamAttr(mapper::DetectionOutputParam *param, const caffe::DetectionOutputParameter &iter) { - if (iter.has_top_k()) { - param->topK = iter.top_k(); - } - if (iter.has_background_label_id()) { - param->backgroundLabelId = iter.background_label_id(); - } - if (iter.has_multi_class_sorting()) { - param->multiClassSorting = iter.multi_class_sorting(); - } - if (iter.has_share_location()) { - param->shareLocation = iter.share_location(); - } - if (iter.has_clip_bbox()) { - param->clipBbox = iter.clip_bbox(); - } - if (iter.has_calc_mode()) { - param->calcMode = iter.calc_mode(); - } - if (iter.has_report_flag()) { - param->reportFlag = iter.report_flag(); - } - if (iter.has_top()) { - param->top = iter.top(); - } - if (iter.has_share_variance()) { - param->shareVariance = iter.share_variance(); - } - if (!iter.variance().empty()) { - param->varianceVec = std::vector(iter.variance().begin(), iter.variance().end()); - } - if (!iter.bias().empty()) { - param->biasVec = std::vector(iter.bias().begin(), iter.bias().end()); - } -} - -ops::PrimitiveC *CaffeDetectionOutputParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeDetectionOutputParser::Parse(const caffe::LayerParameter &proto, + const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -106,46 +33,28 @@ ops::PrimitiveC *CaffeDetectionOutputParser::Parse(const caffe::LayerParameter & prim->set_type("DetectionOutput"); if (proto.has_num_anchors()) { - prim->AddAttr(dpico::kNumAnchors, MakeValue(proto.num_anchors())); + prim->AddAttr(dpico::kNumAnchors, api::MakeValue(proto.num_anchors())); } if (proto.has_num_bboxes_per_grid()) { - prim->AddAttr(dpico::kNumBboxesPerGrid, MakeValue(proto.num_bboxes_per_grid())); + prim->AddAttr(dpico::kNumBboxesPerGrid, api::MakeValue(proto.num_bboxes_per_grid())); } if (proto.has_num_coords()) { - prim->AddAttr(dpico::kNumCoords, MakeValue(proto.num_coords())); + prim->AddAttr(dpico::kNumCoords, api::MakeValue(proto.num_coords())); } if (proto.has_num_classes()) { - prim->AddAttr(dpico::kNumClasses, MakeValue(proto.num_classes())); + prim->AddAttr(dpico::kNumClasses, api::MakeValue(proto.num_classes())); } if (proto.has_num_grids_height()) { - prim->AddAttr(dpico::kNumGridsHeight, MakeValue(proto.num_grids_height())); + prim->AddAttr(dpico::kNumGridsHeight, api::MakeValue(proto.num_grids_height())); } if (proto.has_num_grids_width()) { - prim->AddAttr(dpico::kNumGridsWidth, MakeValue(proto.num_grids_width())); + prim->AddAttr(dpico::kNumGridsWidth, api::MakeValue(proto.num_grids_width())); } - - const auto &detection_output_param = proto.detection_output_param(); - std::vector detect_output_param_vec; - for (const auto &iter : detection_output_param) { - mapper::DetectionOutputParam param; - if (SetParamType(¶m, iter.param_type()) != RET_OK) { - MS_LOG(ERROR) << "Set param type failed."; - return nullptr; - } - if (SetCodeType(¶m, iter.code_type()) != RET_OK) { - MS_LOG(ERROR) << "Set code type failed."; - return nullptr; - } - (void)SetParamAttr(¶m, iter); - auto param_hold_ptr = std::make_shared(param); - if (param_hold_ptr == nullptr) { - MS_LOG(ERROR) << "new DetectionOutputParamHolder failed."; - return nullptr; - } - detect_output_param_vec.push_back(param_hold_ptr); + if (dpico::SetAttrsByDetectionOutputParam(prim, proto) != RET_OK) { + MS_LOG(ERROR) << "set attrs by detection output param failed."; + return nullptr; } - prim->AddAttr(dpico::kDetectionOutputParam, MakeValue(detect_output_param_vec)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeDetectionOutputParser("DetectionOutput", new CaffeDetectionOutputParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.h index 2bfb865d0f..a4a2afc838 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_detection_output_parser.h @@ -29,11 +29,8 @@ class CaffeDetectionOutputParser : public CaffeNodeParser { CaffeDetectionOutputParser() : CaffeNodeParser("detection_output") {} ~CaffeDetectionOutputParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; -STATUS SetParamType(mapper::DetectionOutputParam *param, const caffe::DetectionOutputParameter_ParamType ¶m_type); -STATUS SetCodeType(mapper::DetectionOutputParam *param, const caffe::CodeType &code_type); -void SetParamAttr(mapper::DetectionOutputParam *param, const caffe::DetectionOutputParameter &iter); } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.cc index f453fa64c6..fd8b6ce8ed 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.cc @@ -23,6 +23,7 @@ #include "ops/custom.h" #include "common/op_attr.h" #include "common/check_base.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { @@ -38,44 +39,44 @@ bool IsEltwiseOp(const caffe::EltwiseParameter &eltwiseParam) { std::fabs(eltwiseParam.coeff(0) - 1) <= std::numeric_limits::epsilon() && std::fabs(eltwiseParam.coeff(1) - 1) <= std::numeric_limits::epsilon()); } -int SetEltwiseMode(const caffe::EltwiseParameter &eltwiseParam, ops::PrimitiveC *prim) { +int SetEltwiseMode(const caffe::EltwiseParameter &eltwiseParam, const BaseOperatorPtr &prim) { MS_CHECK_TRUE_MSG(prim != nullptr, RET_ERROR, "prim is nullptr."); if (eltwiseParam.has_operation()) { switch (eltwiseParam.operation()) { case caffe::EltwiseParameter::PROD: - (void)prim->AddAttr(ops::kMode, MakeValue(static_cast(mindspore::EltwiseMode::PROD))); + (void)prim->AddAttr(ops::kMode, api::MakeValue(static_cast(mindspore::EltwiseMode::PROD))); break; case caffe::EltwiseParameter::SUM: - (void)prim->AddAttr(ops::kMode, MakeValue(static_cast(mindspore::EltwiseMode::SUM))); + (void)prim->AddAttr(ops::kMode, api::MakeValue(static_cast(mindspore::EltwiseMode::SUM))); break; case caffe::EltwiseParameter::MAX: - (void)prim->AddAttr(ops::kMode, MakeValue(static_cast(mindspore::EltwiseMode::MAXIMUM))); + (void)prim->AddAttr(ops::kMode, api::MakeValue(static_cast(mindspore::EltwiseMode::MAXIMUM))); break; default: MS_LOG(ERROR) << "Eltwise parse params fail, unsupported operation: " << eltwiseParam.operation(); return RET_ERROR; } } else { - (void)prim->AddAttr(ops::kMode, MakeValue(static_cast(mindspore::EltwiseMode::SUM))); + (void)prim->AddAttr(ops::kMode, api::MakeValue(static_cast(mindspore::EltwiseMode::SUM))); } return RET_OK; } -ops::PrimitiveC *ParseToCustomOp(const caffe::EltwiseParameter &eltwiseParam) { - auto prim = std::make_unique(); +BaseOperatorPtr ParseToCustomOp(const caffe::EltwiseParameter &eltwiseParam) { + auto prim = std::make_shared(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "prim is nullptr."); prim->set_type("Eltwise"); if (eltwiseParam.coeff_size() != 0) { auto coeff_vals = std::vector(eltwiseParam.coeff().begin(), eltwiseParam.coeff().end()); - (void)prim->AddAttr(dpico::kCoeffs, MakeValue(coeff_vals)); + (void)prim->AddAttr(dpico::kCoeffs, api::MakeValue(coeff_vals)); } - if (SetEltwiseMode(eltwiseParam, prim.get()) != RET_OK) { + if (SetEltwiseMode(eltwiseParam, prim) != RET_OK) { MS_LOG(ERROR) << "set eltwise mode failed."; return nullptr; } - return prim.release(); + return prim; } } // namespace -ops::PrimitiveC *CaffeEltwiseParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +BaseOperatorPtr CaffeEltwiseParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { const caffe::EltwiseParameter &eltwiseParam = proto.eltwise_param(); if (eltwiseParam.coeff_size() != 0 && eltwiseParam.coeff_size() != proto.bottom_size()) { MS_LOG(ERROR) << "Coeff size(" << eltwiseParam.coeff_size() @@ -88,17 +89,17 @@ ops::PrimitiveC *CaffeEltwiseParser::Parse(const caffe::LayerParameter &proto, c return nullptr; } else if (proto.bottom_size() == kNums2) { if (IsSubOp(eltwiseParam)) { - auto prim = std::make_unique(); + auto prim = std::make_shared(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "prim is nullptr."); - return prim.release(); + return prim; } else if (IsEltwiseOp(eltwiseParam)) { - auto prim = std::make_unique(); + auto prim = std::make_shared(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "prim is nullptr."); - if (SetEltwiseMode(eltwiseParam, prim.get()) != RET_OK) { + if (SetEltwiseMode(eltwiseParam, prim) != RET_OK) { MS_LOG(ERROR) << "set eltwise mode failed. " << proto.name(); return nullptr; } - return prim.release(); + return prim; } else { return ParseToCustomOp(eltwiseParam); } diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.h index f646ecc4ab..d48c9a62f8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_eltwise_parser.h @@ -27,7 +27,7 @@ class CaffeEltwiseParser : public CaffeNodeParser { CaffeEltwiseParser() : CaffeNodeParser("eltwise") {} ~CaffeEltwiseParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.cc index 42ebf2e32e..ab1b2a4988 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.cc @@ -21,8 +21,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeExpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeExpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -44,7 +44,7 @@ ops::PrimitiveC *CaffeExpParser::Parse(const caffe::LayerParameter &proto, const prim->set_shift(0); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeExpParser("Exp", new CaffeExpParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.h index 0f95976efe..8babc3237d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_exp_parser.h @@ -28,7 +28,7 @@ class CaffeExpParser : public CaffeNodeParser { CaffeExpParser() : CaffeNodeParser("exp") {} ~CaffeExpParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.cc index f62da85016..1a5de5103b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.cc @@ -19,8 +19,11 @@ #include #include #include +#include "common/check_base.h" #include "common/op_attr.h" #include "ops/custom.h" +#include "ops/op_name.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace lite { @@ -35,7 +38,7 @@ STATUS SetExtractAxis(const caffe::ExtractParameter &extract_param, ops::Custom } else if (extract_param.has_slice_dim()) { axis = static_cast(extract_param.slice_dim()); } - prim->AddAttr(ops::kAxis, MakeValue(axis)); + prim->AddAttr(ops::kAxis, api::MakeValue(axis)); std::vector axis_attr(sizeof(int32_t)); if (memcpy_s(axis_attr.data(), axis_attr.size() * sizeof(uint8_t), &axis, sizeof(int32_t)) != EOK) { @@ -54,7 +57,7 @@ STATUS SetExtractSlicePointBegin(const caffe::ExtractParameter &extract_param, o if (extract_param.has_slice_point_begin()) { slice_point_begin = extract_param.slice_point_begin(); } - prim->AddAttr(dpico::kSlicePointBegin, MakeValue(slice_point_begin)); + prim->AddAttr(dpico::kSlicePointBegin, api::MakeValue(slice_point_begin)); std::vector slice_point_begin_attr(sizeof(uint32_t)); if (memcpy_s(slice_point_begin_attr.data(), slice_point_begin_attr.size() * sizeof(uint8_t), &slice_point_begin, @@ -74,7 +77,7 @@ STATUS SetExtractSlicePointEnd(const caffe::ExtractParameter &extract_param, ops if (extract_param.has_slice_point_end()) { slice_point_end = extract_param.slice_point_end(); } - prim->AddAttr(dpico::kSlicePointEnd, MakeValue(slice_point_end)); + prim->AddAttr(dpico::kSlicePointEnd, api::MakeValue(slice_point_end)); std::vector slice_point_end_attr(sizeof(uint32_t)); if (memcpy_s(slice_point_end_attr.data(), slice_point_end_attr.size() * sizeof(uint8_t), &slice_point_end, @@ -86,8 +89,8 @@ STATUS SetExtractSlicePointEnd(const caffe::ExtractParameter &extract_param, ops return RET_OK; } } // namespace -ops::PrimitiveC *CaffeExtractParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeExtractParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -110,7 +113,7 @@ ops::PrimitiveC *CaffeExtractParser::Parse(const caffe::LayerParameter &proto, c } prim->set_attr(custom_attrs); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeExtractParser("Extract", new CaffeExtractParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.h index 108f1ff5d3..ad01b6f5c9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_extract_parser.h @@ -28,7 +28,7 @@ class CaffeExtractParser : public CaffeNodeParser { CaffeExtractParser() : CaffeNodeParser("extract") {} ~CaffeExtractParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.cc index fd0b528c3b..756c7a2d55 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.cc @@ -21,8 +21,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeFlattenParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeFlattenParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -30,14 +30,14 @@ ops::PrimitiveC *CaffeFlattenParser::Parse(const caffe::LayerParameter &proto, c const caffe::FlattenParameter &flatten_param = proto.flatten_param(); if (flatten_param.has_axis()) { - prim->AddAttr(dpico::kStartAxis, MakeValue(flatten_param.axis())); + prim->AddAttr(dpico::kStartAxis, api::MakeValue(flatten_param.axis())); } if (flatten_param.has_end_axis()) { - prim->AddAttr(dpico::kEndAxis, MakeValue(flatten_param.end_axis())); + prim->AddAttr(dpico::kEndAxis, api::MakeValue(flatten_param.end_axis())); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_CaffeFlattenParser("Flatten", new CaffeFlattenParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.h index fa73cbe555..5551d346a0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_flatten_parser.h @@ -27,7 +27,7 @@ class CaffeFlattenParser : public CaffeNodeParser { CaffeFlattenParser() : CaffeNodeParser("flatten") {} ~CaffeFlattenParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace mindspore::lite diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.cc index d35a9ea0ca..5cf1aa7d81 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.cc @@ -23,8 +23,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeGruParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeGruParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -34,26 +34,26 @@ ops::PrimitiveC *CaffeGruParser::Parse(const caffe::LayerParameter &proto, const if (proto.has_recurrent_param()) { const auto &gru_param = proto.recurrent_param(); if (gru_param.has_num_output()) { - prim->AddAttr(dpico::kNumOutput, MakeValue(gru_param.num_output())); + prim->AddAttr(dpico::kNumOutput, api::MakeValue(gru_param.num_output())); } if (gru_param.has_expose_hidden()) { - prim->AddAttr(dpico::kExposeHidden, MakeValue(gru_param.expose_hidden())); - prim->AddAttr(dpico::kOutputLastFrameFlag, MakeValue(gru_param.expose_hidden())); - prim->AddAttr(dpico::kInitialHOnlineFlag, MakeValue(gru_param.expose_hidden())); - prim->AddAttr(dpico::kUseDefaultInitialHFlag, MakeValue(!gru_param.expose_hidden())); + prim->AddAttr(dpico::kExposeHidden, api::MakeValue(gru_param.expose_hidden())); + prim->AddAttr(dpico::kOutputLastFrameFlag, api::MakeValue(gru_param.expose_hidden())); + prim->AddAttr(dpico::kInitialHOnlineFlag, api::MakeValue(gru_param.expose_hidden())); + prim->AddAttr(dpico::kUseDefaultInitialHFlag, api::MakeValue(!gru_param.expose_hidden())); } else { - prim->AddAttr(dpico::kUseDefaultInitialHFlag, MakeValue(true)); + prim->AddAttr(dpico::kUseDefaultInitialHFlag, api::MakeValue(true)); } } // set default value - prim->AddAttr(dpico::kHasSplitHWeightFlag, MakeValue(true)); - prim->AddAttr(dpico::kHasSplitBiasFlag, MakeValue(false)); - prim->AddAttr(dpico::kGruWeightOrderZrhFlag, MakeValue(false)); - prim->AddAttr(dpico::kOnnxModeOutFlag, MakeValue(false)); - prim->AddAttr(dpico::kKeepDirectionDimFlag, MakeValue(false)); + prim->AddAttr(dpico::kHasSplitHWeightFlag, api::MakeValue(true)); + prim->AddAttr(dpico::kHasSplitBiasFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kGruWeightOrderZrhFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kOnnxModeOutFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kKeepDirectionDimFlag, api::MakeValue(false)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeGruParser("GRU", new CaffeGruParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.h index ca16fc4bf4..1eeb88f213 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_gru_parser.h @@ -28,7 +28,7 @@ class CaffeGruParser : public CaffeNodeParser { CaffeGruParser() : CaffeNodeParser("Gru") {} ~CaffeGruParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.cc index be6e58405d..87a500e3e5 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.cc @@ -32,9 +32,9 @@ void TransformShape(caffe::BlobShape *shape) { shape->add_dim(origin_row); } } // namespace -ops::PrimitiveC *CaffeInnerProductParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeInnerProductParser::Parse(const caffe::LayerParameter &proto, + const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -78,7 +78,7 @@ ops::PrimitiveC *CaffeInnerProductParser::Parse(const caffe::LayerParameter &pro MS_LOG(ERROR) << "InnerProduct Parse num_output for " << proto.name().c_str() << " failed."; return nullptr; } else { - prim->AddAttr(dpico::kNumOutput, MakeValue(static_cast(innerProductParam.num_output()))); + prim->AddAttr(dpico::kNumOutput, api::MakeValue(static_cast(innerProductParam.num_output()))); } if (innerProductParam.axis() == 1 || innerProductParam.axis() == dpico::kAxis2) { @@ -92,7 +92,7 @@ ops::PrimitiveC *CaffeInnerProductParser::Parse(const caffe::LayerParameter &pro prim->set_has_bias(true); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeInnerProductParser("InnerProduct", new CaffeInnerProductParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.h index 084a60d694..594bf8046e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_innerproduct_parser.h @@ -28,7 +28,7 @@ class CaffeInnerProductParser : public CaffeNodeParser { CaffeInnerProductParser() : CaffeNodeParser("innerproduct") {} ~CaffeInnerProductParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.cc index 9f70daa7b5..a4a3ae83b5 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.cc @@ -36,17 +36,16 @@ STATUS CaffeInspector::InspectModel(const caffe::NetParameter &proto) { return RET_OK; } -STATUS CaffeInspector::ParseInput() { +void CaffeInspector::ParseInput() { if (net.input_size() > 0) { MS_LOG(INFO) << "This net exist input."; for (int i = 0; i < net.input_size(); i++) { graphInput.insert(net.input(i)); } } - return RET_OK; } -STATUS CaffeInspector::FindGraphInputsAndOutputs() { +void CaffeInspector::FindGraphInputsAndOutputs() { for (const auto &iter : layerBottoms) { if (layerTops.find(iter) == layerTops.end()) { graphInput.insert(iter); @@ -66,10 +65,9 @@ STATUS CaffeInspector::FindGraphInputsAndOutputs() { graphOutput.insert(iter); } } - return RET_OK; } -STATUS CaffeInspector::SetLayerTopsAndBottoms() { +void CaffeInspector::SetLayerTopsAndBottoms() { for (int32_t i = 0; i < net.layer_size(); i++) { auto &layer = const_cast(net.layer(i)); if (layer.top_size() == 1 && layer.bottom_size() == 1 && layer.top(0) == layer.bottom(0)) { @@ -85,7 +83,6 @@ STATUS CaffeInspector::SetLayerTopsAndBottoms() { layerBottoms.insert(layer.bottom(j)); } } - return RET_OK; } } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.h index a6d522a26f..70f1d5b7b3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_inspector.h @@ -32,9 +32,9 @@ class CaffeInspector { ~CaffeInspector() = default; STATUS InspectModel(const caffe::NetParameter &proto); - STATUS ParseInput(); - STATUS FindGraphInputsAndOutputs(); - STATUS SetLayerTopsAndBottoms(); + void ParseInput(); + void FindGraphInputsAndOutputs(); + void SetLayerTopsAndBottoms(); std::set GetGraphInput() { return graphInput; } std::set GetGraphOutput() { return graphOutput; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.cc index 7829976ca6..1d5af40942 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.cc @@ -18,13 +18,13 @@ #include #include "include/registry/converter_context.h" #include "common/op_attr.h" -#include "ops/op_utils.h" #include "ops/resize.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeInterpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeInterpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -52,23 +52,23 @@ ops::PrimitiveC *CaffeInterpParser::Parse(const caffe::LayerParameter &proto, co } if (interp_param.has_zoom_factor()) { - prim->AddAttr(dpico::kZoomFactor, MakeValue(interp_param.zoom_factor())); + prim->AddAttr(dpico::kZoomFactor, api::MakeValue(interp_param.zoom_factor())); } if (interp_param.has_shrink_factor()) { - prim->AddAttr(dpico::kShrinkFactor, MakeValue(interp_param.shrink_factor())); + prim->AddAttr(dpico::kShrinkFactor, api::MakeValue(interp_param.shrink_factor())); } if (interp_param.has_pad_beg()) { - prim->AddAttr(dpico::kPadBeg, MakeValue(interp_param.pad_beg())); + prim->AddAttr(dpico::kPadBeg, api::MakeValue(interp_param.pad_beg())); } if (interp_param.has_pad_end()) { - prim->AddAttr(dpico::kPadEnd, MakeValue(interp_param.pad_end())); + prim->AddAttr(dpico::kPadEnd, api::MakeValue(interp_param.pad_end())); } int fmk_type = converter::FmkType::kFmkTypeCaffe; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim->AddAttr(ops::kFmkType, api::MakeValue(static_cast(fmk_type))); + return prim; } CaffeNodeRegistrar g_caffeInterpParser("Interp", new CaffeInterpParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.h index 1ff1722cbb..e3c0130e10 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_interp_parser.h @@ -28,7 +28,7 @@ class CaffeInterpParser : public CaffeNodeParser { CaffeInterpParser() : CaffeNodeParser("Interp") {} ~CaffeInterpParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.cc index 7dd92db19b..ce07088c52 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.cc @@ -17,11 +17,12 @@ #include "parser/caffe/caffe_log_parser.h" #include #include "ops/log.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeLogParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeLogParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -30,19 +31,19 @@ ops::PrimitiveC *CaffeLogParser::Parse(const caffe::LayerParameter &proto, const const caffe::LogParameter &log_param = proto.log_param(); if (log_param.has_base()) { - prim->AddAttr(ops::kBase, MakeValue(log_param.base())); + prim->AddAttr(ops::kBase, api::MakeValue(log_param.base())); } if (log_param.has_scale()) { - prim->AddAttr(ops::kScale, MakeValue(log_param.scale())); + prim->AddAttr(ops::kScale, api::MakeValue(log_param.scale())); } if (log_param.has_shift()) { - prim->AddAttr(ops::kShift, MakeValue(log_param.shift())); + prim->AddAttr(ops::kShift, api::MakeValue(log_param.shift())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeLogParser("Log", new CaffeLogParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.h index c4823affad..f6aacb2458 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_log_parser.h @@ -28,7 +28,7 @@ class CaffeLogParser : public CaffeNodeParser { CaffeLogParser() : CaffeNodeParser("Log") {} ~CaffeLogParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.cc index 21ac3de956..5272e7c817 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.cc @@ -17,13 +17,13 @@ #include "parser/caffe/caffe_lrn_parser.h" #include #include "common/op_attr.h" -#include "ops/op_utils.h" #include "ops/lrn.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeLRNParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeLRNParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -49,9 +49,9 @@ ops::PrimitiveC *CaffeLRNParser::Parse(const caffe::LayerParameter &proto, const } if (lrnParam.has_norm_region()) { if (lrnParam.norm_region() == caffe::LRNParameter_NormRegion::LRNParameter_NormRegion_WITHIN_CHANNEL) { - prim->AddAttr(ops::kNormRegion, MakeValue("WITHIN_CHANNEL")); + prim->AddAttr(ops::kNormRegion, api::MakeValue("WITHIN_CHANNEL")); } else if (lrnParam.norm_region() == caffe::LRNParameter_NormRegion::LRNParameter_NormRegion_ACROSS_CHANNELS) { - prim->AddAttr(ops::kNormRegion, MakeValue("ACROSS_CHANNELS")); + prim->AddAttr(ops::kNormRegion, api::MakeValue("ACROSS_CHANNELS")); } else { MS_LOG(ERROR) << "invalid norm region param. " << lrnParam.norm_region(); return nullptr; @@ -69,8 +69,8 @@ ops::PrimitiveC *CaffeLRNParser::Parse(const caffe::LayerParameter &proto, const prim->set_alpha(alpha); int two_sides = 2; prim->set_depth_radius(size / two_sides); - prim->AddAttr(dpico::kLrnK, MakeValue(k)); - return prim.release(); + prim->AddAttr(dpico::kLrnK, api::MakeValue(k)); + return prim; } CaffeNodeRegistrar g_caffeLRNParser("LRN", new CaffeLRNParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.h index 62a4cfe9d8..e4c5ffdfa7 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lrn_parser.h @@ -28,7 +28,7 @@ class CaffeLRNParser : public CaffeNodeParser { CaffeLRNParser() : CaffeNodeParser("LRN") {} ~CaffeLRNParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.cc index 9c992d1901..4e79e007c4 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.cc @@ -17,14 +17,13 @@ #include "parser/caffe/caffe_lstm_parser.h" #include #include -#include #include "ops/custom.h" #include "common/op_attr.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeLstmParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeLstmParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -34,25 +33,25 @@ ops::PrimitiveC *CaffeLstmParser::Parse(const caffe::LayerParameter &proto, cons if (proto.has_recurrent_param()) { const auto &lstm_param = proto.recurrent_param(); if (lstm_param.has_num_output()) { - prim->AddAttr(dpico::kNumOutput, MakeValue(lstm_param.num_output())); + prim->AddAttr(dpico::kNumOutput, api::MakeValue(lstm_param.num_output())); } if (lstm_param.has_expose_hidden()) { - prim->AddAttr(dpico::kExposeHidden, MakeValue(lstm_param.expose_hidden())); - prim->AddAttr(dpico::kOutputLastFrameFlag, MakeValue(lstm_param.expose_hidden())); - prim->AddAttr(dpico::kInitialHOnlineFlag, MakeValue(lstm_param.expose_hidden())); - prim->AddAttr(dpico::kUseDefaultInitialHFlag, MakeValue(!lstm_param.expose_hidden())); - prim->AddAttr(dpico::kInitialCOnlineFlag, MakeValue(lstm_param.expose_hidden())); - prim->AddAttr(dpico::kUseDefaultInitialCFlag, MakeValue(!lstm_param.expose_hidden())); + prim->AddAttr(dpico::kExposeHidden, api::MakeValue(lstm_param.expose_hidden())); + prim->AddAttr(dpico::kOutputLastFrameFlag, api::MakeValue(lstm_param.expose_hidden())); + prim->AddAttr(dpico::kInitialHOnlineFlag, api::MakeValue(lstm_param.expose_hidden())); + prim->AddAttr(dpico::kUseDefaultInitialHFlag, api::MakeValue(!lstm_param.expose_hidden())); + prim->AddAttr(dpico::kInitialCOnlineFlag, api::MakeValue(lstm_param.expose_hidden())); + prim->AddAttr(dpico::kUseDefaultInitialCFlag, api::MakeValue(!lstm_param.expose_hidden())); } } // set default value - prim->AddAttr(dpico::kKeepDirectionDimFlag, MakeValue(false)); - prim->AddAttr(dpico::kPeepHoleFlag, MakeValue(false)); - prim->AddAttr(dpico::kLstmWeightOrderIofcFlag, MakeValue(false)); - prim->AddAttr(dpico::kSequenceLensOnlineFlag, MakeValue(false)); + prim->AddAttr(dpico::kKeepDirectionDimFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kPeepHoleFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kLstmWeightOrderIofcFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kSequenceLensOnlineFlag, api::MakeValue(false)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeLSTMParser("LSTM", new CaffeLstmParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.h index 235e819128..919e2e73a7 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_lstm_parser.h @@ -28,7 +28,7 @@ class CaffeLstmParser : public CaffeNodeParser { CaffeLstmParser() : CaffeNodeParser("Lstm") {} ~CaffeLstmParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.cc index 0524d9e85e..a02be39b1f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.cc @@ -21,8 +21,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeMatmulParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeMatmulParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -30,19 +30,19 @@ ops::PrimitiveC *CaffeMatmulParser::Parse(const caffe::LayerParameter &proto, co if (proto.has_matmul_param()) { const caffe::MatMulParameter &matmul_param = proto.matmul_param(); if (matmul_param.has_dim_1()) { - prim->AddAttr(dpico::kDim1, MakeValue(matmul_param.dim_1())); + prim->AddAttr(dpico::kDim1, api::MakeValue(matmul_param.dim_1())); } if (matmul_param.has_dim_2()) { - prim->AddAttr(dpico::kDim2, MakeValue(matmul_param.dim_2())); + prim->AddAttr(dpico::kDim2, api::MakeValue(matmul_param.dim_2())); } if (matmul_param.has_dim_3()) { - prim->AddAttr(dpico::kDim3, MakeValue(matmul_param.dim_3())); + prim->AddAttr(dpico::kDim3, api::MakeValue(matmul_param.dim_3())); } } prim->set_transpose_a(false); prim->set_transpose_b(true); prim->set_activation_type(mindspore::NO_ACTIVATION); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeMatmulParser("MatMul", new CaffeMatmulParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.h index ba997758e6..b97c1ce3fb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_matmul_parser.h @@ -28,7 +28,7 @@ class CaffeMatmulParser : public CaffeNodeParser { CaffeMatmulParser() : CaffeNodeParser("Matmul") {} ~CaffeMatmulParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.cc index 8f2922bd91..8a9dbf2c63 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.cc @@ -19,13 +19,14 @@ #include #include #include +#include #include "parser/caffe/caffe_inspector.h" #include "parser/caffe/caffe_node_parser_registry.h" #include "common/anf_util.h" #include "parser/parser_utils.h" #include "common/op_enum.h" #include "parser/unify_format.h" -#include "api/ir/func_graph.h" +#include "mindapi/ir/func_graph.h" #include "include/registry/converter_context.h" #include "ops/make_tuple.h" #include "ops/return.h" @@ -93,8 +94,8 @@ api::FuncGraphPtr CaffeModelParser::Parse(const converter::ConverterParameters & MS_LOG(ERROR) << "convert graph outputs failed."; return nullptr; } - res_graph_->set_attr("graph_name", MakeValue("main_graph")); - res_graph_->set_attr("fmk", MakeValue(static_cast(kFmkTypeCaffe))); + res_graph_->set_attr("graph_name", api::MakeValue("main_graph")); + res_graph_->set_attr("fmk", api::MakeValue(static_cast(kFmkTypeCaffe))); std::set all_func_graphs = {}; GetAllFuncGraph(res_graph_, &all_func_graphs); if (PostAdjust(all_func_graphs) != RET_OK) { @@ -157,15 +158,15 @@ STATUS CaffeModelParser::ConvertLayers() { continue; } - auto primitive_c = node_parser->Parse(layer, weight); - if (primitive_c == nullptr) { + auto base_operator_ptr = node_parser->Parse(layer, weight); + if (base_operator_ptr == nullptr) { MS_LOG(ERROR) << "parse node " << layer.name() << " failed."; status = RET_ERROR; continue; } // build inputs - std::vector input_nodes; + std::vector input_nodes; status = ConvertBottom(layer, &input_nodes); if (status != RET_OK) { MS_LOG(ERROR) << "Convert layer bottom for " << layer.name() << " failed."; @@ -173,7 +174,7 @@ STATUS CaffeModelParser::ConvertLayers() { } // build weights - std::vector const_parameters; + std::vector const_parameters; status = ConvertBlobs(weight, &const_parameters); if (status != RET_OK) { MS_LOG(ERROR) << "Convert blobs for " << layer.name() << " failed."; @@ -181,7 +182,8 @@ STATUS CaffeModelParser::ConvertLayers() { } // build cnode - std::vector op_inputs = {NewValueNode(std::shared_ptr(primitive_c))}; + api::SharedPtr primitive(std::move(base_operator_ptr)); + std::vector op_inputs = {api::NewValueNode(primitive)}; op_inputs.insert(op_inputs.end(), input_nodes.begin(), input_nodes.end()); op_inputs.insert(op_inputs.end(), const_parameters.begin(), const_parameters.end()); auto new_cnode = res_graph_->NewCNode(op_inputs); @@ -310,13 +312,13 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { CaffeInspector caffeInspector; caffeInspector.InspectModel(caffe_model_); if (caffeInspector.GetGraphOutput().size() > 1) { - std::vector make_tuple_inputs; - auto make_tuple_prim_ptr = std::make_shared(); + std::vector make_tuple_inputs; + auto make_tuple_prim_ptr = api::MakeShared(); if (make_tuple_prim_ptr == nullptr) { MS_LOG(ERROR) << "new MakeTuple failed"; return RET_NULL_PTR; } - auto make_tuple_prim = NewValueNode(make_tuple_prim_ptr); + auto make_tuple_prim = api::NewValueNode(make_tuple_prim_ptr); make_tuple_inputs.emplace_back(make_tuple_prim); for (const auto &output_node : caffeInspector.GetGraphOutput()) { if (nodes_.find(output_node) == nodes_.end()) { @@ -329,26 +331,26 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { auto make_tuple_cnode = res_graph_->NewCNode(make_tuple_inputs); make_tuple_cnode->set_fullname_with_scope("return tuple"); - std::vector op_inputs; - auto return_prim_ptr = std::make_shared(); + std::vector op_inputs; + auto return_prim_ptr = api::MakeShared(); if (return_prim_ptr == nullptr) { MS_LOG(ERROR) << "new Return failed"; return RET_NULL_PTR; } - auto value_node = NewValueNode(return_prim_ptr); + auto value_node = api::NewValueNode(return_prim_ptr); op_inputs.emplace_back(value_node); op_inputs.emplace_back(make_tuple_cnode); auto cnode = res_graph_->NewCNode(op_inputs); cnode->set_fullname_with_scope("Return"); res_graph_->set_return(cnode); } else { - auto returnPrim = std::make_shared(); - if (returnPrim == nullptr) { + auto return_prim = api::MakeShared(); + if (return_prim == nullptr) { MS_LOG(ERROR) << "new Return failed"; return RET_NULL_PTR; } - auto valueNode = NewValueNode(returnPrim); - std::vector opInputs{valueNode}; + auto valueNode = api::NewValueNode(return_prim); + std::vector opInputs{valueNode}; if (nodes_.find(*caffeInspector.GetGraphOutput().begin()) == nodes_.end()) { MS_LOG(ERROR) << "Can't find input node."; return RET_NOT_FIND_OP; @@ -366,7 +368,8 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { return RET_OK; } -STATUS CaffeModelParser::ConvertBlobs(const caffe::LayerParameter &layer, std::vector *const_parameters) { +STATUS CaffeModelParser::ConvertBlobs(const caffe::LayerParameter &layer, + std::vector *const_parameters) { if (const_parameters == nullptr) { MS_LOG(ERROR) << "const parameters are null"; return RET_NULL_PTR; @@ -385,7 +388,6 @@ STATUS CaffeModelParser::ConvertBlobs(const caffe::LayerParameter &layer, std::v // cal Weight num auto parameter = res_graph_->add_parameter(); - auto type_ptr = TypeIdToType(TypeId::kNumberTypeFloat32); std::vector shape_vector; (void)std::transform(shape.begin(), shape.end(), std::back_inserter(shape_vector), [](const int32_t &value) { return static_cast(value); }); @@ -400,7 +402,7 @@ STATUS CaffeModelParser::ConvertBlobs(const caffe::LayerParameter &layer, std::v } int count = 0; - tensor::TensorPtr tensor_info = nullptr; + api::TensorPtr tensor_info = nullptr; if (layer.blobs(i).double_data_size() > 0) { count = layer.blobs(i).double_data_size(); auto buf = std::make_unique(count); @@ -435,7 +437,7 @@ STATUS CaffeModelParser::ConvertBlobs(const caffe::LayerParameter &layer, std::v return RET_OK; } -STATUS CaffeModelParser::ConvertBottom(const caffe::LayerParameter &layer, std::vector *input_nodes) { +STATUS CaffeModelParser::ConvertBottom(const caffe::LayerParameter &layer, std::vector *input_nodes) { if (input_nodes == nullptr) { MS_LOG(ERROR) << "input_nodes is null"; return RET_NULL_PTR; @@ -461,7 +463,7 @@ STATUS CaffeModelParser::ConvertBottom(const caffe::LayerParameter &layer, std:: return RET_OK; } -STATUS CaffeModelParser::ConvertTop(const caffe::LayerParameter &layer, const CNodePtr &cnode) { +STATUS CaffeModelParser::ConvertTop(const caffe::LayerParameter &layer, const api::CNodePtr &cnode) { if (layer.top_size() == 1) { auto abstract = dpico::CreateTensorAbstract({}, kNumberTypeFloat32); if (abstract == nullptr) { @@ -473,7 +475,7 @@ STATUS CaffeModelParser::ConvertTop(const caffe::LayerParameter &layer, const CN return RET_OK; } - AbstractBasePtrList abstract_list; + api::AbstractBasePtrList abstract_list; for (int i = 0; i < layer.top_size(); i++) { auto abstract = dpico::CreateTensorAbstract({}, kNumberTypeFloat32); if (abstract == nullptr) { @@ -481,19 +483,19 @@ STATUS CaffeModelParser::ConvertTop(const caffe::LayerParameter &layer, const CN return RET_ERROR; } abstract_list.emplace_back(abstract); - auto tuple_get_item_prim_ptr = std::make_shared(); + auto tuple_get_item_prim_ptr = api::MakeShared(); if (tuple_get_item_prim_ptr == nullptr) { MS_LOG(ERROR) << "new TupleGetItem failed"; return RET_NULL_PTR; } - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); - auto get_item_value = NewValueNode(MakeValue(i)); - std::vector inputs{tuple_get_item_prim, cnode, get_item_value}; - CNodePtr get_item_cnode = res_graph_->NewCNode(inputs); + auto tuple_get_item_prim = api::NewValueNode(tuple_get_item_prim_ptr); + auto get_item_value = api::NewValueNode(api::MakeValue(i)); + std::vector inputs{tuple_get_item_prim, cnode, get_item_value}; + api::CNodePtr get_item_cnode = res_graph_->NewCNode(inputs); get_item_cnode->set_fullname_with_scope(layer.top(i)); nodes_[layer.top(i)] = get_item_cnode; } - auto abstract_tuple = std::make_shared(abstract_list); + auto abstract_tuple = api::MakeShared(abstract_list); if (abstract_tuple == nullptr) { MS_LOG(ERROR) << "abstract_tuple is nullptr."; return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.h index c25755093e..c917eb2a03 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_model_parser.h @@ -25,7 +25,6 @@ #include "include/registry/model_parser.h" #include "include/registry/model_parser_registry.h" #include "./pico_caffe.pb.h" -#include "ops/primitive_c.h" using STATUS = int; namespace mindspore::lite { @@ -48,18 +47,18 @@ class CaffeModelParser : public converter::ModelParser { STATUS ConvertLayers(); - STATUS ConvertBlobs(const caffe::LayerParameter &layer, std::vector *const_parameters); + STATUS ConvertBlobs(const caffe::LayerParameter &layer, std::vector *const_parameters); - STATUS ConvertBottom(const caffe::LayerParameter &layer, std::vector *input_nodes); + STATUS ConvertBottom(const caffe::LayerParameter &layer, std::vector *input_nodes); - STATUS ConvertTop(const caffe::LayerParameter &layer, const CNodePtr &cnode); + STATUS ConvertTop(const caffe::LayerParameter &layer, const api::CNodePtr &cnode); std::string GetOriginLayerName(const std::string &layer_name); caffe::NetParameter caffe_model_; caffe::NetParameter caffe_weight_; std::unordered_map caffe_layers_; - std::unordered_map nodes_; + std::unordered_map nodes_; }; } // namespace mindspore::lite diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.cc index 588276ee1e..938b341a53 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.cc @@ -19,11 +19,12 @@ #include #include "common/op_attr.h" #include "ops/custom.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeMvnParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeMvnParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -33,17 +34,17 @@ ops::PrimitiveC *CaffeMvnParser::Parse(const caffe::LayerParameter &proto, const if (proto.has_mvn_param()) { const caffe::MVNParameter &mvn_parameter = proto.mvn_param(); if (mvn_parameter.has_eps()) { - prim->AddAttr(ops::kEps, MakeValue(mvn_parameter.eps())); + prim->AddAttr(ops::kEps, api::MakeValue(mvn_parameter.eps())); } if (mvn_parameter.has_across_channels()) { - prim->AddAttr(dpico::kAcrossChannels, MakeValue(mvn_parameter.across_channels())); + prim->AddAttr(dpico::kAcrossChannels, api::MakeValue(mvn_parameter.across_channels())); } if (mvn_parameter.has_normalize_variance()) { - prim->AddAttr(dpico::kNormalizeVariance, MakeValue(mvn_parameter.normalize_variance())); + prim->AddAttr(dpico::kNormalizeVariance, api::MakeValue(mvn_parameter.normalize_variance())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeMvnParser("MVN", new CaffeMvnParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.h index 2fd79685b2..30e84ebbf7 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_mvn_parser.h @@ -28,7 +28,7 @@ class CaffeMvnParser : public CaffeNodeParser { CaffeMvnParser() : CaffeNodeParser("Mvn") {} ~CaffeMvnParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_node_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_node_parser.h index 1c5142c0c7..3afd341536 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_node_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_node_parser.h @@ -21,16 +21,17 @@ #include #include #include +#include #include "google/protobuf/message.h" #include "./pico_caffe.pb.h" #include "common/anf_util.h" -#include "utils/check_convert_utils.h" -#include "ops/primitive_c.h" +#include "ops/base_operator.h" using mindspore::lite::STATUS; namespace mindspore { namespace lite { +using BaseOperatorPtr = std::shared_ptr; constexpr int kNums2 = 2; constexpr int kNums4 = 4; class CaffeNodeParser { @@ -39,7 +40,7 @@ class CaffeNodeParser { virtual ~CaffeNodeParser() = default; - virtual ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + virtual BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { return nullptr; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.cc index 7871a7225a..07117c7301 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.cc @@ -22,15 +22,15 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeNopParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeNopParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; } prim->set_type("Nop"); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeNopParser("Nop", new CaffeNopParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.h index 6ccb4c2a60..d38cc0c8a1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_nop_parser.h @@ -28,7 +28,7 @@ class CaffeNopParser : public CaffeNodeParser { CaffeNopParser() : CaffeNodeParser("Nop") {} ~CaffeNopParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.cc index 926a924aff..c94d94cc86 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.cc @@ -19,11 +19,12 @@ #include #include "common/op_attr.h" #include "ops/custom.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeNormalizeParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeNormalizeParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -33,20 +34,20 @@ ops::PrimitiveC *CaffeNormalizeParser::Parse(const caffe::LayerParameter &proto, if (proto.has_norm_param()) { const caffe::NormalizeParameter &normalize_param = proto.norm_param(); if (normalize_param.has_across_spatial()) { - prim->AddAttr(dpico::kAcrossSpatial, MakeValue(normalize_param.across_spatial())); + prim->AddAttr(dpico::kAcrossSpatial, api::MakeValue(normalize_param.across_spatial())); } if (normalize_param.has_channel_shared()) { - prim->AddAttr(dpico::kChannelShared, MakeValue(normalize_param.channel_shared())); + prim->AddAttr(dpico::kChannelShared, api::MakeValue(normalize_param.channel_shared())); } if (normalize_param.has_sqrt_a()) { - prim->AddAttr(dpico::kSqrtA, MakeValue(normalize_param.sqrt_a())); + prim->AddAttr(dpico::kSqrtA, api::MakeValue(normalize_param.sqrt_a())); } if (normalize_param.has_eps()) { - prim->AddAttr(ops::kEps, MakeValue(normalize_param.eps())); + prim->AddAttr(ops::kEps, api::MakeValue(normalize_param.eps())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeNormalizeParser("Normalize", new CaffeNormalizeParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.h index 28d61634dd..e5f773f464 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_normalize_parser.h @@ -28,7 +28,7 @@ class CaffeNormalizeParser : public CaffeNodeParser { CaffeNormalizeParser() : CaffeNodeParser("Normalize") {} ~CaffeNormalizeParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.cc index e69a76397b..80f80ee375 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.cc @@ -20,12 +20,12 @@ #include #include "ops/custom.h" #include "common/op_attr.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePassThroughParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffePassThroughParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -37,7 +37,7 @@ ops::PrimitiveC *CaffePassThroughParser::Parse(const caffe::LayerParameter &prot const auto &pass_through_param = proto.pass_through_param(); if (pass_through_param.has_num_output()) { uint32_t num_output = pass_through_param.num_output(); - prim->AddAttr(dpico::kNumOutput, MakeValue(num_output)); // for mapper + prim->AddAttr(dpico::kNumOutput, api::MakeValue(num_output)); // for mapper std::vector num_output_attr(sizeof(uint32_t)); if (memcpy_s(num_output_attr.data(), num_output_attr.size() * sizeof(uint8_t), &num_output, sizeof(uint32_t)) != @@ -49,7 +49,7 @@ ops::PrimitiveC *CaffePassThroughParser::Parse(const caffe::LayerParameter &prot } if (pass_through_param.has_block_height()) { uint32_t block_height = pass_through_param.block_height(); - prim->AddAttr(dpico::kBlockHeight, MakeValue(block_height)); + prim->AddAttr(dpico::kBlockHeight, api::MakeValue(block_height)); std::vector block_height_attr(sizeof(uint32_t)); if (memcpy_s(block_height_attr.data(), block_height_attr.size() * sizeof(uint8_t), &block_height, @@ -61,7 +61,7 @@ ops::PrimitiveC *CaffePassThroughParser::Parse(const caffe::LayerParameter &prot } if (pass_through_param.has_block_width()) { uint32_t block_width = pass_through_param.block_width(); - prim->AddAttr(dpico::kBlockWidth, MakeValue(block_width)); + prim->AddAttr(dpico::kBlockWidth, api::MakeValue(block_width)); std::vector block_width_attr(sizeof(uint32_t)); if (memcpy_s(block_width_attr.data(), block_width_attr.size() * sizeof(uint8_t), &block_width, @@ -74,7 +74,7 @@ ops::PrimitiveC *CaffePassThroughParser::Parse(const caffe::LayerParameter &prot } prim->set_attr(custom_attrs); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffePassThroughParser("PassThrough", new CaffePassThroughParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.h index 906631ed60..29f47cc1ce 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pass_through_parser.h @@ -28,7 +28,7 @@ class CaffePassThroughParser : public CaffeNodeParser { CaffePassThroughParser() : CaffeNodeParser("PassThrough") {} ~CaffePassThroughParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.cc index a599a38dc6..f1ab4f187c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.cc @@ -21,8 +21,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePermuteParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffePermuteParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -34,9 +34,9 @@ ops::PrimitiveC *CaffePermuteParser::Parse(const caffe::LayerParameter &proto, c for (int i = 0; i < num_order_dims; ++i) { perm[i] = permuteParam.order()[i]; } - prim->AddAttr(dpico::kPerm, MakeValue(perm)); + prim->AddAttr(dpico::kPerm, api::MakeValue(perm)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffePermuteParser("Permute", new CaffePermuteParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.h index 04d5789b7b..b9c4c29361 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_permute_parser.h @@ -28,7 +28,7 @@ class CaffePermuteParser : public CaffeNodeParser { CaffePermuteParser() : CaffeNodeParser("Permute") {} ~CaffePermuteParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.cc index ba0c73565b..8e076f441d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.cc @@ -99,7 +99,7 @@ mindspore::RoundMode CaffePoolingParser::ParseRoundMode(const caffe::PoolingPara return roundMode; } -ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +BaseOperatorPtr CaffePoolingParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { const caffe::PoolingParameter &poolingParam = proto.pooling_param(); // parse kernel params @@ -127,7 +127,7 @@ ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, c auto roundMode = ParseRoundMode(poolingParam); if (poolingParam.pool() == caffe::PoolingParameter::MAX) { - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -139,9 +139,9 @@ ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, c prim->set_pad(pad); prim->set_round_mode(roundMode); prim->set_global(poolingParam.global_pooling()); - return prim.release(); + return prim; } else if (poolingParam.pool() == caffe::PoolingParameter::AVE) { - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -153,7 +153,7 @@ ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, c prim->set_pad(pad); prim->set_round_mode(roundMode); prim->set_global(poolingParam.global_pooling()); - return prim.release(); + return prim; } else { MS_LOG(ERROR) << "poolingParam.pool() is not MAX or AVE"; return nullptr; diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.h index e73e9e0568..3914dacc56 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_pooling_parser.h @@ -20,6 +20,7 @@ #include #include "parser/caffe/caffe_node_parser.h" #include "parser/caffe/caffe_node_parser_registry.h" +#include "mindapi/base/types.h" namespace mindspore { namespace lite { @@ -28,7 +29,7 @@ class CaffePoolingParser : public CaffeNodeParser { CaffePoolingParser() : CaffeNodeParser("pooling") {} ~CaffePoolingParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; static STATUS ParsePads(const caffe::PoolingParameter &poolingParam, std::vector *pad); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.cc index 9bcc70fde6..57e06a6c4c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.cc @@ -18,11 +18,12 @@ #include #include #include "ops/fusion/pow_fusion.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePowerParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffePowerParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -42,11 +43,11 @@ ops::PrimitiveC *CaffePowerParser::Parse(const caffe::LayerParameter &proto, con shift = powerParam.shift(); } } - prim->AddAttr(ops::kPower, MakeValue(power)); + prim->AddAttr(ops::kPower, api::MakeValue(power)); prim->set_scale(scale); prim->set_shift(shift); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffePowerParser("Power", new CaffePowerParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.h index 9af3ea5b3b..f60a9c10eb 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_power_parser.h @@ -28,7 +28,7 @@ class CaffePowerParser : public CaffeNodeParser { CaffePowerParser() : CaffeNodeParser("power") {} ~CaffePowerParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.cc index 88f2cfce6e..caeecbec60 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.cc @@ -5,7 +5,7 @@ * 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.0f + * 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, @@ -20,8 +20,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffePReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -33,7 +33,7 @@ ops::PrimitiveC *CaffePReluParser::Parse(const caffe::LayerParameter &proto, con prim->set_channel_shared(false); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffePReluParser("PReLU", new CaffePReluParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.h index 626ca369ff..afba90f384 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_prelu_parser.h @@ -28,7 +28,7 @@ class CaffePReluParser : public CaffeNodeParser { CaffePReluParser() : CaffeNodeParser("pRelu") {} ~CaffePReluParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.cc index bf9edd2680..9595a765bc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.cc @@ -20,12 +20,13 @@ #include #include "common/op_attr.h" #include "ops/custom.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePSROIPoolingParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffePSROIPoolingParser::Parse(const caffe::LayerParameter &proto, + const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -37,7 +38,7 @@ ops::PrimitiveC *CaffePSROIPoolingParser::Parse(const caffe::LayerParameter &pro const auto &psroi_pooling_param = proto.psroi_pooling_param(); if (psroi_pooling_param.has_spatial_scale()) { float spatial_scale = psroi_pooling_param.spatial_scale(); - prim->AddAttr(dpico::kSpatialScale, MakeValue(spatial_scale)); + prim->AddAttr(dpico::kSpatialScale, api::MakeValue(spatial_scale)); std::vector spatial_scale_attr(sizeof(float)); if (memcpy_s(spatial_scale_attr.data(), spatial_scale_attr.size() * sizeof(uint8_t), &spatial_scale, @@ -49,7 +50,7 @@ ops::PrimitiveC *CaffePSROIPoolingParser::Parse(const caffe::LayerParameter &pro } if (psroi_pooling_param.has_group_size()) { int32_t group_size = psroi_pooling_param.group_size(); - prim->AddAttr(dpico::kGroupSize, MakeValue(group_size)); + prim->AddAttr(dpico::kGroupSize, api::MakeValue(group_size)); std::vector group_size_attr(sizeof(int32_t)); if (memcpy_s(group_size_attr.data(), group_size_attr.size() * sizeof(uint8_t), &group_size, sizeof(int32_t)) != @@ -61,7 +62,7 @@ ops::PrimitiveC *CaffePSROIPoolingParser::Parse(const caffe::LayerParameter &pro } if (psroi_pooling_param.has_output_dim()) { int32_t output_dim = psroi_pooling_param.output_dim(); - prim->AddAttr(dpico::kOutputDim, MakeValue(output_dim)); + prim->AddAttr(dpico::kOutputDim, api::MakeValue(output_dim)); std::vector output_dim_attr(sizeof(int32_t)); if (memcpy_s(output_dim_attr.data(), output_dim_attr.size() * sizeof(uint8_t), &output_dim, sizeof(int32_t)) != @@ -74,7 +75,7 @@ ops::PrimitiveC *CaffePSROIPoolingParser::Parse(const caffe::LayerParameter &pro } prim->set_attr(custom_attrs); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffePSROIPoolingParser("PSROIPooling", new CaffePSROIPoolingParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.h index cf1955c0df..91c6222cb6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_psroi_pooling_parser.h @@ -28,7 +28,7 @@ class CaffePSROIPoolingParser : public CaffeNodeParser { CaffePSROIPoolingParser() : CaffeNodeParser("PSROIPooling") {} ~CaffePSROIPoolingParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.cc index 6cef4a517e..0bf2879011 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.cc @@ -17,14 +17,14 @@ #include "parser/caffe/caffe_reduce_parser.h" #include #include -#include "ops/op_utils.h" #include "include/registry/converter_context.h" #include "ops/fusion/reduce_fusion.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReduceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeReduceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -56,14 +56,14 @@ ops::PrimitiveC *CaffeReduceParser::Parse(const caffe::LayerParameter &proto, co } else { axes = std::vector(1, 0); } - prim->AddAttr(ops::kAxes, MakeValue(axes)); + prim->AddAttr(ops::kAxes, api::MakeValue(axes)); if (reduce_param.has_coeff()) { - prim->AddAttr(ops::kCoeff, MakeValue(reduce_param.coeff())); + prim->AddAttr(ops::kCoeff, api::MakeValue(reduce_param.coeff())); } int fmk_type = converter::FmkType::kFmkTypeCaffe; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim->AddAttr(ops::kFmkType, api::MakeValue(static_cast(fmk_type))); + return prim; } CaffeNodeRegistrar g_caffeReduceParser("Reduction", new CaffeReduceParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.h index c079a41a32..af68ece5a0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reduce_parser.h @@ -28,7 +28,7 @@ class CaffeReduceParser : public CaffeNodeParser { CaffeReduceParser() : CaffeNodeParser("reduce") {} ~CaffeReduceParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.cc index 13c21ce4bf..22d192081b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.cc @@ -17,13 +17,13 @@ #include "parser/caffe/caffe_reshape_parser.h" #include #include "common/op_attr.h" -#include "ops/op_utils.h" #include "ops/reshape.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReshapeParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeReshapeParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -38,17 +38,17 @@ ops::PrimitiveC *CaffeReshapeParser::Parse(const caffe::LayerParameter &proto, c for (int i = 0; i < blob_shape.dim_size(); i++) { shape.push_back(blob_shape.dim(i)); } - prim->AddAttr(ops::kShape, MakeValue(shape)); + prim->AddAttr(ops::kShape, api::MakeValue(shape)); if (reshapeParam.has_axis()) { - prim->AddAttr(ops::kAxis, MakeValue(static_cast(reshapeParam.axis()))); + prim->AddAttr(ops::kAxis, api::MakeValue(reshapeParam.axis())); } if (reshapeParam.has_num_axes()) { - prim->AddAttr(dpico::kNumAxes, MakeValue(static_cast(reshapeParam.num_axes()))); + prim->AddAttr(dpico::kNumAxes, api::MakeValue(reshapeParam.num_axes())); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeReshapeParser("Reshape", new CaffeReshapeParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.h index f226384657..4f270dd140 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reshape_parser.h @@ -28,7 +28,7 @@ class CaffeReshapeParser : public CaffeNodeParser { CaffeReshapeParser() : CaffeNodeParser("reshape") {} ~CaffeReshapeParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.cc index af312e486a..6a77124855 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.cc @@ -20,8 +20,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReverseParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeReverseParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -33,7 +33,7 @@ ops::PrimitiveC *CaffeReverseParser::Parse(const caffe::LayerParameter &proto, c prim->set_axis({0}); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeReverseParser("Reverse", new CaffeReverseParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.h index 432aefc36a..d1956d4769 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_reverse_parser.h @@ -28,7 +28,7 @@ class CaffeReverseParser : public CaffeNodeParser { CaffeReverseParser() : CaffeNodeParser("reverse") {} ~CaffeReverseParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.cc index 2722635012..1df08ef90d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.cc @@ -23,8 +23,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeRnnParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeRnnParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -34,21 +34,21 @@ ops::PrimitiveC *CaffeRnnParser::Parse(const caffe::LayerParameter &proto, const if (proto.has_recurrent_param()) { const auto &rnn_param = proto.recurrent_param(); if (rnn_param.has_num_output()) { - prim->AddAttr(dpico::kNumOutput, MakeValue(rnn_param.num_output())); + prim->AddAttr(dpico::kNumOutput, api::MakeValue(rnn_param.num_output())); } if (rnn_param.has_expose_hidden()) { - prim->AddAttr(dpico::kExposeHidden, MakeValue(rnn_param.expose_hidden())); - prim->AddAttr(dpico::kOutputLastFrameFlag, MakeValue(rnn_param.expose_hidden())); - prim->AddAttr(dpico::kInitialHOnlineFlag, MakeValue(rnn_param.expose_hidden())); - prim->AddAttr(dpico::kUseDefaultInitialHFlag, MakeValue(!rnn_param.expose_hidden())); + prim->AddAttr(dpico::kExposeHidden, api::MakeValue(rnn_param.expose_hidden())); + prim->AddAttr(dpico::kOutputLastFrameFlag, api::MakeValue(rnn_param.expose_hidden())); + prim->AddAttr(dpico::kInitialHOnlineFlag, api::MakeValue(rnn_param.expose_hidden())); + prim->AddAttr(dpico::kUseDefaultInitialHFlag, api::MakeValue(!rnn_param.expose_hidden())); } } // set default value - prim->AddAttr(dpico::kKeepDirectionDimFlag, MakeValue(false)); - prim->AddAttr(dpico::kHasOutputGateFlag, MakeValue(true)); + prim->AddAttr(dpico::kKeepDirectionDimFlag, api::MakeValue(false)); + prim->AddAttr(dpico::kHasOutputGateFlag, api::MakeValue(true)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeRnnParser("RNN", new CaffeRnnParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.h index 844041f483..4320cc834f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_rnn_parser.h @@ -28,7 +28,7 @@ class CaffeRnnParser : public CaffeNodeParser { CaffeRnnParser() : CaffeNodeParser("Rnn") {} ~CaffeRnnParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.cc index d4ac4052f1..4d7eb11e3f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.cc @@ -20,8 +20,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeROIPoolingParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeROIPoolingParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -39,7 +39,7 @@ ops::PrimitiveC *CaffeROIPoolingParser::Parse(const caffe::LayerParameter &proto } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeROIPoolingParser("ROIPooling", new CaffeROIPoolingParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.h index 15db7da0a2..96a9f16898 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_roi_pooling_parser.h @@ -28,7 +28,7 @@ class CaffeROIPoolingParser : public CaffeNodeParser { CaffeROIPoolingParser() : CaffeNodeParser("ROIPooling") {} ~CaffeROIPoolingParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.cc index 287e697ad9..4e8d08fa7c 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.cc @@ -36,8 +36,8 @@ STATUS CaffeScaleParser::GetAxisIndex(const int32_t &axis, uint32_t *axis_index) return RET_OK; } -ops::PrimitiveC *CaffeScaleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeScaleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -60,13 +60,13 @@ ops::PrimitiveC *CaffeScaleParser::Parse(const caffe::LayerParameter &proto, con } if (scaleParam.has_bias_term()) { - (void)prim->AddAttr(dpico::kBiasTerm, MakeValue(scaleParam.bias_term())); + (void)prim->AddAttr(dpico::kBiasTerm, api::MakeValue(scaleParam.bias_term())); } if (scaleParam.has_num_axes()) { - (void)prim->AddAttr(dpico::kNumAxes, MakeValue(scaleParam.num_axes())); + (void)prim->AddAttr(dpico::kNumAxes, api::MakeValue(scaleParam.num_axes())); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeScaleParser("Scale", new CaffeScaleParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.h index f4fff2f128..0f220bf5c3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_scale_parser.h @@ -28,7 +28,7 @@ class CaffeScaleParser : public CaffeNodeParser { CaffeScaleParser() : CaffeNodeParser("scale") {} ~CaffeScaleParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; private: STATUS GetAxisIndex(const int32_t &axis, uint32_t *axis_index); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.cc index 4cd226d6d2..a0e2b45fd1 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.cc @@ -17,14 +17,14 @@ #include "parser/caffe/caffe_shuffle_channel_parser.h" #include #include -#include "common/op_attr.h" #include "ops/custom.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeShuffleChannelParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeShuffleChannelParser::Parse(const caffe::LayerParameter &proto, + const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -34,11 +34,11 @@ ops::PrimitiveC *CaffeShuffleChannelParser::Parse(const caffe::LayerParameter &p if (proto.has_shuffle_channel_param()) { const caffe::ShuffleChannelParameter &shuffle_channel_parameter = proto.shuffle_channel_param(); if (shuffle_channel_parameter.has_group()) { - prim->AddAttr(ops::kGroup, MakeValue(shuffle_channel_parameter.group())); + prim->AddAttr(ops::kGroup, api::MakeValue(shuffle_channel_parameter.group())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeShuffleChannelParser("ShuffleChannel", new CaffeShuffleChannelParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.h index 0202367ca9..c3fb95c613 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_shuffle_channel_parser.h @@ -28,7 +28,7 @@ class CaffeShuffleChannelParser : public CaffeNodeParser { CaffeShuffleChannelParser() : CaffeNodeParser("ShuffleChannel") {} ~CaffeShuffleChannelParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.cc index 1a11531d9c..274b0a7daf 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.cc @@ -23,8 +23,8 @@ namespace lite { namespace { const int kDefaultAxis = 1; } // namespace -ops::PrimitiveC *CaffeSliceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeSliceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -53,7 +53,7 @@ ops::PrimitiveC *CaffeSliceParser::Parse(const caffe::LayerParameter &proto, con prim->set_axis(kDefaultAxis); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeSliceParser("Slice", new CaffeSliceParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.h index 6482ff414f..09eba2a69b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_slice_parser.h @@ -28,7 +28,7 @@ class CaffeSliceParser : public CaffeNodeParser { CaffeSliceParser() : CaffeNodeParser("slice") {} ~CaffeSliceParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.cc index d978c80b9b..d5c0849232 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.cc @@ -20,8 +20,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeSoftmaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeSoftmaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -35,7 +35,7 @@ ops::PrimitiveC *CaffeSoftmaxParser::Parse(const caffe::LayerParameter &proto, c prim->set_axis({1}); } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeSoftmaxParser("Softmax", new CaffeSoftmaxParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.h index 7cd66e3da3..0ca0d2a98d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_softmax_parser.h @@ -28,7 +28,7 @@ class CaffeSoftmaxParser : public CaffeNodeParser { CaffeSoftmaxParser() : CaffeNodeParser("softmax") {} ~CaffeSoftmaxParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.cc index 0ca08eee16..0aa0cfd33f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.cc @@ -21,26 +21,27 @@ #include #include "common/op_attr.h" #include "ops/custom.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeSppParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeSppParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; } prim->set_type("Spp"); - int pool_method = 0; + int64_t pool_method = 0; if (proto.has_spp_param()) { const caffe::SPPParameter &spp_param = proto.spp_param(); if (spp_param.has_pool()) { - pool_method = static_cast(spp_param.pool()); + pool_method = static_cast(spp_param.pool()); } if (spp_param.has_pyramid_height()) { - auto pyramid_height = spp_param.pyramid_height(); - prim->AddAttr(dpico::kPyramidHeight, MakeValue(pyramid_height)); + uint32_t pyramid_height = spp_param.pyramid_height(); + prim->AddAttr(dpico::kPyramidHeight, api::MakeValue(static_cast(pyramid_height))); std::map> custom_attrs; std::vector pyramid_height_attr(sizeof(uint32_t)); if (memcpy_s(pyramid_height_attr.data(), pyramid_height_attr.size() * sizeof(uint8_t), &pyramid_height, @@ -55,9 +56,9 @@ ops::PrimitiveC *CaffeSppParser::Parse(const caffe::LayerParameter &proto, const return nullptr; } } - prim->AddAttr(dpico::kPoolMethod, MakeValue(pool_method)); + prim->AddAttr(dpico::kPoolMethod, api::MakeValue(pool_method)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeSppParser("SPP", new CaffeSppParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.h index c8bf2ed823..b481fec1b8 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_spp_parser.h @@ -28,7 +28,7 @@ class CaffeSppParser : public CaffeNodeParser { CaffeSppParser() : CaffeNodeParser("Spp") {} ~CaffeSppParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.cc index 87d5a59e6a..3853574c0f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.cc @@ -22,8 +22,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeThresholdParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeThresholdParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -33,11 +33,11 @@ ops::PrimitiveC *CaffeThresholdParser::Parse(const caffe::LayerParameter &proto, if (proto.has_threshold_param()) { const caffe::ThresholdParameter &threshold_param = proto.threshold_param(); if (threshold_param.has_threshold()) { - prim->AddAttr(dpico::kThreshold, MakeValue(threshold_param.threshold())); + prim->AddAttr(dpico::kThreshold, api::MakeValue(threshold_param.threshold())); } } - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeThresholdParser("Threshold", new CaffeThresholdParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.h index 8982e70b1b..d597879a6a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_threshold_parser.h @@ -28,7 +28,7 @@ class CaffeThresholdParser : public CaffeNodeParser { CaffeThresholdParser() : CaffeNodeParser("Threshold") {} ~CaffeThresholdParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.cc index 555f76d1d1..a4d9bbd003 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.cc @@ -22,8 +22,8 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeTileParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +BaseOperatorPtr CaffeTileParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -45,9 +45,9 @@ ops::PrimitiveC *CaffeTileParser::Parse(const caffe::LayerParameter &proto, cons } else { multiples.push_back(1); } - prim->AddAttr(dpico::kMultiples, MakeValue(multiples)); + prim->AddAttr(dpico::kMultiples, api::MakeValue(multiples)); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeTileParser("Tile", new CaffeTileParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.h index 966328ebb6..bf26ad28c9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_tile_parser.h @@ -28,7 +28,7 @@ class CaffeTileParser : public CaffeNodeParser { CaffeTileParser() : CaffeNodeParser("tile") {} ~CaffeTileParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.cc index a72ad9c9ca..338c24ac62 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.cc @@ -18,17 +18,18 @@ #include #include #include -#include "./op_enum_public.h" #include "common/op_attr.h" #include "ops/custom.h" +#include "ops/op_name.h" +#include "third_party/securec/include/securec.h" namespace mindspore { namespace lite { namespace { constexpr float kDefaultScaleVal = 2.0; -} -ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { - auto prim = std::make_unique(); +} // namespace +BaseOperatorPtr CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr."; return nullptr; @@ -44,7 +45,7 @@ ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, } if (upsample_param.upsample_h()) { uint32_t upsample_h = upsample_param.upsample_h(); - prim->AddAttr(dpico::kUpsampleH, MakeValue(upsample_h)); + prim->AddAttr(dpico::kUpsampleH, api::MakeValue(upsample_h)); std::vector upsample_h_attr(sizeof(uint32_t)); if (memcpy_s(upsample_h_attr.data(), upsample_h_attr.size() * sizeof(uint8_t), &upsample_h, sizeof(uint32_t)) != @@ -56,7 +57,7 @@ ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, } if (upsample_param.upsample_w()) { uint32_t upsample_w = upsample_param.upsample_w(); - prim->AddAttr(dpico::kUpsampleH, MakeValue(upsample_w)); + prim->AddAttr(dpico::kUpsampleH, api::MakeValue(upsample_w)); std::vector upsample_w_attr(sizeof(uint32_t)); if (memcpy_s(upsample_w_attr.data(), upsample_w_attr.size() * sizeof(uint8_t), &upsample_w, sizeof(uint32_t)) != @@ -70,10 +71,10 @@ ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, auto mode = upsample_param.interpolation_mode(); switch (mode) { case caffe::UpsampleParameter_InterpolationMode_NEAREST: - prim->AddAttr(dpico::kInterpolationMode, MakeValue(dpico::kNearest)); + prim->AddAttr(dpico::kInterpolationMode, api::MakeValue(dpico::kNearest)); break; case caffe::UpsampleParameter_InterpolationMode_BILINEAR: - prim->AddAttr(dpico::kInterpolationMode, MakeValue(dpico::kBilinear)); + prim->AddAttr(dpico::kInterpolationMode, api::MakeValue(dpico::kBilinear)); break; default: MS_LOG(ERROR) << "current interpolation mode is not supported. " << mode; @@ -82,7 +83,7 @@ ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, } } - prim->AddAttr(ops::kScale, MakeValue(scale)); + prim->AddAttr(ops::kScale, api::MakeValue(scale)); std::vector scale_attr(sizeof(float)); if (memcpy_s(scale_attr.data(), scale_attr.size() * sizeof(uint8_t), &scale, sizeof(float)) != EOK) { MS_LOG(ERROR) << "memcpy_s failed."; @@ -90,7 +91,7 @@ ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, } custom_attrs[ops::kScale] = scale_attr; prim->set_attr(custom_attrs); - return prim.release(); + return prim; } CaffeNodeRegistrar g_caffeUpsampleParser("Upsample", new CaffeUpsampleParser()); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.h index 4cb2952bb5..38a1d4d914 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/caffe_upsample_parser.h @@ -28,7 +28,7 @@ class CaffeUpsampleParser : public CaffeNodeParser { CaffeUpsampleParser() : CaffeNodeParser("Upsample") {} ~CaffeUpsampleParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + BaseOperatorPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.cc index d4e6a4362a..041b6041ec 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.cc @@ -18,6 +18,17 @@ #include #include "common/op_enum.h" #include "common/anf_util.h" +#include "common/check_base.h" +#include "ops/transpose.h" +#include "ops/reshape.h" +#include "ops/gather.h" +#include "ops/cast.h" +#include "ops/fusion/topk_fusion.h" +#include "ops/fusion/tile_fusion.h" +#include "ops/fusion/reduce_fusion.h" +#include "ops/fusion/pad_fusion.h" +#include "ops/fusion/pow_fusion.h" +#include "ops/resize.h" namespace mindspore::lite { namespace { @@ -25,29 +36,29 @@ constexpr int kBuildInputFlagTwo = 2; constexpr int kBuildInputFlagThree = 3; constexpr int kBuildInputFlagFour = 4; } // namespace -STATUS InputAdjust::AddAttrToInput(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, int input_num, +STATUS InputAdjust::AddAttrToInput(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, int input_num, const std::string &attr_name, int flag) { - MS_ASSERT(cnode != nullptr); + MS_CHECK_TRUE_MSG(cnode != nullptr, RET_ERROR, "cnode is nullptr"); if (!dpico::CheckInputs(cnode)) { MS_LOG(ERROR) << "input is invalid."; return lite::RET_INPUT_TENSOR_ERROR; } - auto primitive_c = GetValueNode(cnode->input(0)); - auto value_ptr = primitive_c->GetAttr(attr_name); + auto primitive = api::GetValueNode(cnode->input(0)); + auto value_ptr = primitive->GetAttr(attr_name); if (value_ptr == nullptr) { MS_LOG(DEBUG) << "there is no attr :" << attr_name; return lite::RET_NO_CHANGE; } auto inputs = cnode->inputs(); if (static_cast(inputs.size()) > input_num) { - primitive_c->EraseAttr(attr_name); + primitive->EraseAttr(attr_name); MS_LOG(DEBUG) << "input num has been meet, which is " << inputs.size(); return lite::RET_OK; } else if (static_cast(inputs.size()) < input_num) { MS_LOG(ERROR) << "input num is invalid."; return lite::RET_ERROR; } - AnfNodePtr param_node; + api::AnfNodePtr param_node; switch (flag) { case 1: { auto value_data_vec = dpico::CastToInt(value_ptr); @@ -71,7 +82,7 @@ STATUS InputAdjust::AddAttrToInput(const api::FuncGraphPtr &func_graph, const CN break; } case kBuildInputFlagFour: { - auto value_data = GetValue(value_ptr); + auto value_data = api::GetValue(value_ptr); param_node = dpico::BuildFloatValueParameterNode(func_graph, value_data, cnode->fullname_with_scope() + "_" + attr_name); break; @@ -81,7 +92,7 @@ STATUS InputAdjust::AddAttrToInput(const api::FuncGraphPtr &func_graph, const CN return lite::RET_ERROR; } } - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_ASSERT(manager != nullptr); manager->AddEdge(cnode, param_node); @@ -98,38 +109,38 @@ bool InputAdjust::Run(const api::FuncGraphPtr &func_graph) { auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); STATUS status = lite::RET_OK; for (auto &node : node_list) { - auto cnode = node->cast(); + auto cnode = node->cast(); if (cnode == nullptr) { continue; } - if (dpico::CheckPrimitiveType(node, prim::kPrimTranspose)) { + if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust Transpose"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "perm", kBuildInputFlagTwo); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimReshape)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust Reshape"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "shape", kBuildInputFlagTwo); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimGather)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust Gather"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex3, "axis", 1); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimCast)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust Cast"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "to", 1); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimTopKFusion)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust TopKFusion"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "k", 1); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimTileFusion)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust TileFusion"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "multiples", kBuildInputFlagTwo); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimReduceFusion)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust ReduceFusion"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "axes", kBuildInputFlagTwo); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimPadFusion)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust PadFusion"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "paddings", kBuildInputFlagThree); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimPowFusion)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { MS_LOG(INFO) << "Adjust PowFuison"; status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "power", kBuildInputFlagFour); - } else if (dpico::CheckPrimitiveType(node, prim::kPrimResize)) { + } else if (dpico::CheckPrimitiveType(node, api::MakeShared())) { status = AddAttrToInput(func_graph, cnode, dpico::kInputIndex2, "zoom_factor", 1); } if (status != lite::RET_OK && status != lite::RET_NO_CHANGE) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.h b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.h index 94fe510273..d80893205b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/caffe/inputs_adjust.h @@ -20,20 +20,18 @@ #include #include #include -#include "ir/anf.h" -#include "api/ir/func_graph.h" -#include "ops/primitive_c.h" +#include "mindapi/ir/anf.h" +#include "mindapi/ir/func_graph.h" #include "include/errorcode.h" using mindspore::lite::STATUS; -using PrimitiveCPtr = std::shared_ptr; namespace mindspore::lite { class InputAdjust { public: InputAdjust() {} ~InputAdjust() = default; - STATUS AddAttrToInput(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, int input_num, + STATUS AddAttrToInput(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, int input_num, const std::string &attr_name, int flag); bool Run(const api::FuncGraphPtr &func_graph); }; diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_helper.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_helper.cc new file mode 100644 index 0000000000..b6d510a512 --- /dev/null +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_helper.cc @@ -0,0 +1,247 @@ +/** + * Copyright 2022 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 "parser/detection_output_param_helper.h" +#include +#include +#include "mindapi/base/logging.h" +#include "common/op_attr.h" +#include "include/errorcode.h" +#include "ops/custom.h" +#include "./pico_caffe.pb.h" +#include "op/detection_output_operator.h" + +using mindspore::lite::RET_ERROR; +using mindspore::lite::RET_OK; +namespace mindspore { +namespace dpico { +namespace { +int GetProposalParamType(mapper::ProposalParamType *proposal_param_type, + const caffe::DetectionOutputParameter_ParamType ¶m_type) { + if (proposal_param_type == nullptr) { + MS_LOG(ERROR) << "input proposal_param_type is nullptr. "; + return RET_ERROR; + } + switch (param_type) { + case caffe::DetectionOutputParameter_ParamType_DecBBox: + *proposal_param_type = mapper::ProposalParamType::PROPOSAL_DECBBOX; + break; + case caffe::DetectionOutputParameter_ParamType_Sort: + *proposal_param_type = mapper::ProposalParamType::PROPOSAL_SORT; + break; + case caffe::DetectionOutputParameter_ParamType_Nms: + *proposal_param_type = mapper::ProposalParamType::PROPOSAL_NMS; + break; + case caffe::DetectionOutputParameter_ParamType_FilterBox: + *proposal_param_type = mapper::ProposalParamType::PROPOSAL_FILTERBOX; + break; + default: + MS_LOG(ERROR) << "Unsupported Param Type: " << param_type; + return RET_ERROR; + } + return RET_OK; +} +int GetCodeType(mapper::DecBboxCodeType *dec_bbox_code_type, const caffe::CodeType &code_type) { + if (dec_bbox_code_type == nullptr) { + MS_LOG(ERROR) << "input dec_bbox_code_type is nullptr. "; + return RET_ERROR; + } + if (code_type == caffe::CodeType::CENTER_SIZE) { + *dec_bbox_code_type = mapper::DecBboxCodeType::DECBBOX_CODE_TYPE_CENTER_SIZE; + } else { + MS_LOG(ERROR) << "Unsupported Code type: " << code_type; + return RET_ERROR; + } + return RET_OK; +} + +int SetAttrsByParam(const std::shared_ptr &custom_prim, + const caffe::DetectionOutputParameter &detection_param, int index) { + if (detection_param.has_top_k()) { + custom_prim->AddAttr(kDetectionTopK + std::to_string(index), api::MakeValue(detection_param.top_k())); + } + if (detection_param.has_background_label_id()) { + custom_prim->AddAttr(kDetectionBackgroundLabelId + std::to_string(index), + api::MakeValue(detection_param.background_label_id())); + } + if (detection_param.has_multi_class_sorting()) { + custom_prim->AddAttr(kDetectionMultiClassSorting + std::to_string(index), + api::MakeValue(detection_param.multi_class_sorting())); + } + if (detection_param.has_share_location()) { + custom_prim->AddAttr(kDetectionShareLocation + std::to_string(index), + api::MakeValue(detection_param.share_location())); + } + if (detection_param.has_clip_bbox()) { + custom_prim->AddAttr(kDetectionClipBbox + std::to_string(index), api::MakeValue(detection_param.clip_bbox())); + } + if (detection_param.has_calc_mode()) { + custom_prim->AddAttr(kDetectionCalcMode + std::to_string(index), api::MakeValue(detection_param.calc_mode())); + } + if (detection_param.has_report_flag()) { + custom_prim->AddAttr(kDetectionReportFlag + std::to_string(index), + api::MakeValue(detection_param.report_flag())); + } + if (detection_param.has_top()) { + custom_prim->AddAttr(kDetectionTop + std::to_string(index), api::MakeValue(detection_param.top())); + } + if (detection_param.has_share_variance()) { + custom_prim->AddAttr(kDetectionShareVariance + std::to_string(index), + api::MakeValue(detection_param.share_variance())); + } + if (!detection_param.variance().empty()) { + custom_prim->AddAttr(kDetectionVarianceVec + std::to_string(index), + api::MakeValue>( + std::vector(detection_param.variance().begin(), detection_param.variance().end()))); + } + if (!detection_param.bias().empty()) { + custom_prim->AddAttr(kDetectionBiasVec + std::to_string(index), + api::MakeValue>( + std::vector(detection_param.bias().begin(), detection_param.bias().end()))); + } + + if (detection_param.has_param_type()) { + mapper::ProposalParamType proposal_param_type; + if (GetProposalParamType(&proposal_param_type, detection_param.param_type()) != RET_OK) { + MS_LOG(ERROR) << "get detection proposal param type failed."; + return RET_ERROR; + } + custom_prim->AddAttr(kDetectionProposalParamType + std::to_string(index), + api::MakeValue(static_cast(proposal_param_type))); + } + + if (detection_param.has_code_type()) { + mapper::DecBboxCodeType dec_bbox_code_type; + if (GetCodeType(&dec_bbox_code_type, detection_param.code_type()) != RET_OK) { + MS_LOG(ERROR) << "get detection code type failed."; + return RET_ERROR; + } + custom_prim->AddAttr(kDetectionCodeType + std::to_string(index), + api::MakeValue(static_cast(dec_bbox_code_type))); + } + return RET_OK; +} +mapper::DetectionOutputParam GetParamFromAttrs(const api::SharedPtr &custom_prim, int index) { + mapper::DetectionOutputParam detection_output_param; + auto top_k_ptr = custom_prim->GetAttr(kDetectionTopK + std::to_string(index)); + if (top_k_ptr != nullptr) { + detection_output_param.topK = static_cast(api::GetValue(top_k_ptr)); + } + auto background_label_id_ptr = custom_prim->GetAttr(kDetectionBackgroundLabelId + std::to_string(index)); + if (background_label_id_ptr != nullptr) { + detection_output_param.backgroundLabelId = static_cast(api::GetValue(background_label_id_ptr)); + } + auto multi_class_sorting_ptr = custom_prim->GetAttr(kDetectionMultiClassSorting + std::to_string(index)); + if (multi_class_sorting_ptr != nullptr) { + detection_output_param.multiClassSorting = api::GetValue(multi_class_sorting_ptr); + } + auto share_location_ptr = custom_prim->GetAttr(kDetectionShareLocation + std::to_string(index)); + if (share_location_ptr != nullptr) { + detection_output_param.shareLocation = api::GetValue(share_location_ptr); + } + auto clip_bbox_ptr = custom_prim->GetAttr(kDetectionClipBbox + std::to_string(index)); + if (clip_bbox_ptr != nullptr) { + detection_output_param.clipBbox = api::GetValue(clip_bbox_ptr); + } + auto calc_mode_ptr = custom_prim->GetAttr(kDetectionCalcMode + std::to_string(index)); + if (calc_mode_ptr != nullptr) { + detection_output_param.calcMode = static_cast(api::GetValue(calc_mode_ptr)); + } + auto report_flag_ptr = custom_prim->GetAttr(kDetectionReportFlag + std::to_string(index)); + if (report_flag_ptr != nullptr) { + detection_output_param.reportFlag = api::GetValue(report_flag_ptr); + } + auto top_ptr = custom_prim->GetAttr(kDetectionTop + std::to_string(index)); + if (top_ptr != nullptr) { + detection_output_param.top = api::GetValue(top_ptr); + } + auto share_variance_ptr = custom_prim->GetAttr(kDetectionShareVariance + std::to_string(index)); + if (share_variance_ptr != nullptr) { + detection_output_param.shareVariance = api::GetValue(share_variance_ptr); + } + auto variance_vec_ptr = custom_prim->GetAttr(kDetectionVarianceVec + std::to_string(index)); + if (variance_vec_ptr != nullptr) { + detection_output_param.varianceVec = api::GetValue>(variance_vec_ptr); + } + auto bias_vec_ptr = custom_prim->GetAttr(kDetectionBiasVec + std::to_string(index)); + if (bias_vec_ptr != nullptr) { + detection_output_param.biasVec = api::GetValue>(bias_vec_ptr); + } + auto proposal_param_type_ptr = custom_prim->GetAttr(kDetectionProposalParamType + std::to_string(index)); + if (proposal_param_type_ptr != nullptr) { + detection_output_param.paramType = + static_cast(api::GetValue(proposal_param_type_ptr)); + } + auto code_type_ptr = custom_prim->GetAttr(kDetectionCodeType + std::to_string(index)); + if (code_type_ptr != nullptr) { + detection_output_param.codeType = static_cast(api::GetValue(code_type_ptr)); + } + return detection_output_param; +} +} // namespace +int SetAttrsByDetectionOutputParam(const std::shared_ptr &custom_prim, + const caffe::LayerParameter &proto) { + int detection_output_param_size = proto.detection_output_param_size(); + custom_prim->AddAttr(kDetectionOutputParamSize, api::MakeValue(detection_output_param_size)); + if (detection_output_param_size == 0) { + MS_LOG(INFO) << "no detection param found"; + return RET_OK; + } + for (int i = 0; i < proto.detection_output_param_size(); i++) { + const auto &detect_param = proto.detection_output_param(i); + if (SetAttrsByParam(custom_prim, detect_param, i) != RET_OK) { + MS_LOG(ERROR) << "set prim attrs from detection param failed."; + return RET_ERROR; + } + } + return RET_OK; +} +int SetAttrsByDecBboxParam(const std::shared_ptr &custom_prim, const caffe::LayerParameter &proto) { + if (!proto.has_decbbox_param()) { + MS_LOG(INFO) << "no decbbox param found"; + return RET_OK; + } + custom_prim->AddAttr(kDetectionOutputParamSize, api::MakeValue(1)); + const auto &detect_param = proto.decbbox_param(); + if (SetAttrsByParam(custom_prim, detect_param, 0) != RET_OK) { + MS_LOG(ERROR) << "set prim attrs from decbbox param failed."; + return RET_ERROR; + } + return RET_OK; +} +int GetDetectionOutputParamFromAttrs(std::vector *detection_params, + const api::SharedPtr &custom_prim) { + if (detection_params == nullptr) { + MS_LOG(ERROR) << "input detection_params is nullptr."; + return RET_ERROR; + } + int detection_output_param_size = 0; + auto detection_output_param_size_attr = custom_prim->GetAttr(kDetectionOutputParamSize); + if (detection_output_param_size_attr != nullptr) { + detection_output_param_size = static_cast(api::GetValue(detection_output_param_size_attr)); + } + if (detection_output_param_size == 0) { + MS_LOG(INFO) << "no detection param attr found."; + return RET_OK; + } + for (int i = 0; i < detection_output_param_size; i++) { + detection_params->emplace_back(GetParamFromAttrs(custom_prim, i)); + } + return RET_OK; +} + +} // namespace dpico +} // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_helper.h b/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_helper.h new file mode 100644 index 0000000000..5fce48b1a2 --- /dev/null +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_helper.h @@ -0,0 +1,36 @@ +/** + * Copyright 2022 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. + */ + +#ifndef DPICO_PARSER_DETECTION_OUTPUT_PARAM_HELPER_H_ +#define DPICO_PARSER_DETECTION_OUTPUT_PARAM_HELPER_H_ + +#include +#include +#include +#include +#include "op/detection_output_operator.h" +#include "ops/custom.h" +#include "./pico_caffe.pb.h" + +namespace mindspore { +namespace dpico { +int SetAttrsByDetectionOutputParam(const std::shared_ptr &custom_prim, const caffe::LayerParameter &proto); +int SetAttrsByDecBboxParam(const std::shared_ptr &custom_prim, const caffe::LayerParameter &proto); +int GetDetectionOutputParamFromAttrs(std::vector *detection_params, + const api::SharedPtr &custom_prim); +} // namespace dpico +} // namespace mindspore +#endif // DPICO_PARSER_DETECTION_OUTPUT_PARAM_HELPER_H_ diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_holder.h b/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_holder.h deleted file mode 100644 index 81bd476432..0000000000 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/detection_output_param_holder.h +++ /dev/null @@ -1,45 +0,0 @@ -/** - * 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. - */ - -#ifndef DPICO_PARSER_DETECTION_OUTPUT_PARAM_HOLDER_H_ -#define DPICO_PARSER_DETECTION_OUTPUT_PARAM_HOLDER_H_ - -#include -#include -#include "ir/anf.h" -#include "op/detection_output_operator.h" - -namespace mindspore { -namespace lite { -class DetectionOutputParamHolder : public Value { - public: - explicit DetectionOutputParamHolder(mapper::DetectionOutputParam param) : detection_output_param_(std::move(param)) {} - ~DetectionOutputParamHolder() override = default; - - MS_DECLARE_PARENT(DetectionOutputParamHolder, Value); - - bool operator==(const Value &rhs) const override { // unused - return rhs.isa(); - } - const mapper::DetectionOutputParam &GetDetectionOutputParam() const { return detection_output_param_; } - - private: - mapper::DetectionOutputParam detection_output_param_; -}; -using DetectionOutputParamHolderPtr = std::shared_ptr; -} // namespace lite -} // namespace mindspore -#endif // DPICO_PARSER_DETECTION_OUTPUT_PARAM_HOLDER_H_ diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.cc index 970afd3229..7cc8edf34b 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.cc @@ -23,11 +23,14 @@ #include "ops/custom.h" #include "./onnx.pb.h" #include "include/registry/node_parser_registry.h" +#include "mindapi/base/logging.h" +#include "third_party/securec/include/securec.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxRoiAlignParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { - auto prim = std::make_unique(); +ops::BaseOperatorPtr OnnxRoiAlignParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { + auto prim = api::MakeShared(); if (prim == nullptr) { MS_LOG(ERROR) << "new Custom prim failed."; return nullptr; @@ -52,11 +55,11 @@ ops::PrimitiveC *OnnxRoiAlignParser::Parse(const onnx::GraphProto &onnx_graph, c } } // set attr for mapper - prim->AddAttr(ops::kMode, MakeValue(pool_mode)); - prim->AddAttr(dpico::kOutputHeight, MakeValue(output_height)); - prim->AddAttr(dpico::kOutputWidth, MakeValue(output_width)); - prim->AddAttr(dpico::kSamplingRatio, MakeValue(sampling_ratio)); - prim->AddAttr(dpico::kSpatialScale, MakeValue(spatial_scale)); + prim->AddAttr(ops::kMode, api::MakeValue(pool_mode)); + prim->AddAttr(dpico::kOutputHeight, api::MakeValue(output_height)); + prim->AddAttr(dpico::kOutputWidth, api::MakeValue(output_width)); + prim->AddAttr(dpico::kSamplingRatio, api::MakeValue(sampling_ratio)); + prim->AddAttr(dpico::kSpatialScale, api::MakeValue(spatial_scale)); // set attr for infershape std::map> custom_attrs; @@ -77,7 +80,7 @@ ops::PrimitiveC *OnnxRoiAlignParser::Parse(const onnx::GraphProto &onnx_graph, c custom_attrs[dpico::kOutputWidth] = output_width_attr; prim->set_attr(custom_attrs); - return prim.release(); + return prim; } } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.h b/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.h index 9f6e7f7b74..a5ffafa256 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/onnx/onnx_roi_align_parser.h @@ -17,16 +17,19 @@ #ifndef DPICO_PARSER_ONNX_ONNX_ROI_ALIGN_PARSER_H_ #define DPICO_PARSER_ONNX_ONNX_ROI_ALIGN_PARSER_H_ +#include #include "include/registry/node_parser.h" +#include "ops/base_operator.h" namespace mindspore { namespace lite { +using BaseOperatorPtr = std::shared_ptr; class OnnxRoiAlignParser : public converter::NodeParser { public: OnnxRoiAlignParser() : NodeParser() {} ~OnnxRoiAlignParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + ops::BaseOperatorPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.cc index 7555cb36e0..cce0a218dc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.cc @@ -20,20 +20,27 @@ #include #include #include "common/file_util.h" -#include "ops/op_utils.h" #include "common/format_utils.h" #include "common/anf_util.h" #include "parser/caffe/inputs_adjust.h" #include "common/data_transpose_utils.h" +#include "ops/fusion/conv2d_fusion.h" +#include "ops/fusion/conv2d_transpose_fusion.h" +#include "ops/adam.h" +#include "ops/apply_momentum.h" +#include "ops/sgd.h" +#include "ops/op_name.h" +#include "common/check_base.h" namespace mindspore::lite { namespace { const int WARNING_THRESHOLD = 536870912 * 2; -bool IsWeightNodeSensitive(const AnfNodePtr &node) { - return dpico::CheckPrimitiveType(node, prim::kPrimConv2DFusion) || - dpico::CheckPrimitiveType(node, prim::kPrimConv2dTransposeFusion) || - dpico::CheckPrimitiveType(node, prim::kPrimApplyMomentum) || dpico::CheckPrimitiveType(node, prim::kPrimSGD) || - dpico::CheckPrimitiveType(node, prim::kPrimAdam); +bool IsWeightNodeSensitive(const api::AnfNodePtr &node) { + return dpico::CheckPrimitiveType(node, api::MakeShared()) || + dpico::CheckPrimitiveType(node, api::MakeShared()) || + dpico::CheckPrimitiveType(node, api::MakeShared()) || + dpico::CheckPrimitiveType(node, api::MakeShared()) || + dpico::CheckPrimitiveType(node, api::MakeShared()); } int GetTransposePerm(mindspore::Format src_format, mindspore::Format dst_format, std::vector *perm) { @@ -74,9 +81,11 @@ int GetTransposePermSharing(mindspore::Format src_format, mindspore::Format dst_ return lite::RET_OK; } -int UnifyVariableConvWeight(const api::FuncGraphPtr &graph, const AnfNodePtr &weight_node, mindspore::Format src_format, - mindspore::Format dst_format, std::set *has_visited) { - MS_ASSERT(graph != nullptr && weight_node != nullptr && has_visited != nullptr); +int UnifyVariableConvWeight(const api::FuncGraphPtr &graph, const api::AnfNodePtr &weight_node, + mindspore::Format src_format, mindspore::Format dst_format, + std::set *has_visited) { + MS_CHECK_TRUE_MSG(graph != nullptr && weight_node != nullptr && has_visited != nullptr, RET_ERROR, + "input param contains nullptr."); if (src_format == dst_format) { return lite::RET_OK; } @@ -87,13 +96,12 @@ int UnifyVariableConvWeight(const api::FuncGraphPtr &graph, const AnfNodePtr &we return status; } auto manager = api::FuncGraphManager::Manage(graph); - MS_ASSERT(manager != nullptr); - CNodePtr trans_cnode = nullptr; - auto node_map = manager->node_users(); - auto &weight_node_users = node_map[weight_node]; + MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "manager is nullptr."); + api::CNodePtr trans_cnode = nullptr; + auto weight_node_users = manager->GetUsers(weight_node); for (auto &weight_node_user : weight_node_users) { auto post_node = weight_node_user.first; - if (!utils::isa(post_node)) { + if (!api::utils::isa(post_node)) { MS_LOG(ERROR) << "post node is invalid."; return RET_ERROR; } @@ -126,7 +134,7 @@ int UnifyVariableConvWeight(const api::FuncGraphPtr &graph, const AnfNodePtr &we abstract = dpico::CreateTensorAbstract(shape, TypeId::kNumberTypeFloat32); MS_ASSERT(abstract != nullptr); } - auto shape_ptr = std::make_shared(shape); + auto shape_ptr = api::MakeShared(shape); if (shape_ptr == nullptr) { MS_LOG(ERROR) << "shape ptr is nullptr."; return RET_ERROR; @@ -134,16 +142,17 @@ int UnifyVariableConvWeight(const api::FuncGraphPtr &graph, const AnfNodePtr &we abstract->set_shape(shape_ptr); trans_cnode->set_abstract(abstract); } - auto post_cnode = post_node->cast(); + auto post_cnode = post_node->cast(); manager->SetEdge(post_cnode, weight_node_user.second, trans_cnode); } return RET_OK; } -int HandleConstConvWeightShared(const api::FuncGraphPtr &graph, const AnfNodePtr &weight_node, +int HandleConstConvWeightShared(const api::FuncGraphPtr &graph, const api::AnfNodePtr &weight_node, mindspore::Format src_format, mindspore::Format dst_format, - std::set *has_visited) { - MS_ASSERT(graph != nullptr && weight_node != nullptr && has_visited != nullptr); + std::set *has_visited) { + MS_CHECK_TRUE_MSG(graph != nullptr && weight_node != nullptr && has_visited != nullptr, RET_ERROR, + "input param contains nullptr."); if (src_format == dst_format) { return RET_OK; } @@ -154,13 +163,12 @@ int HandleConstConvWeightShared(const api::FuncGraphPtr &graph, const AnfNodePtr return status; } auto manager = api::FuncGraphManager::Manage(graph); - MS_ASSERT(manager != nullptr); - CNodePtr trans_cnode = nullptr; - auto node_map = manager->node_users(); - auto &weight_node_users = node_map[weight_node]; + MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "manager is nullptr."); + api::CNodePtr trans_cnode = nullptr; + auto weight_node_users = manager->GetUsers(weight_node); for (auto &weight_node_user : weight_node_users) { auto post_node = weight_node_user.first; - if (!utils::isa(post_node)) { + if (!api::utils::isa(post_node)) { MS_LOG(ERROR) << "post node is invalid."; return RET_ERROR; } @@ -172,9 +180,9 @@ int HandleConstConvWeightShared(const api::FuncGraphPtr &graph, const AnfNodePtr trans_cnode = dpico::GenTransposeNode(graph, weight_node, perm, weight_node->fullname_with_scope() + "_post_perm"); MS_ASSERT(trans_cnode != nullptr); - auto prim = GetValueNode(trans_cnode->input(0)); + auto prim = api::GetValueNode(trans_cnode->input(0)); MS_ASSERT(prim != nullptr); - prim->AddAttr(ops::kFormat, MakeValue(dst_format)); + prim->AddAttr(ops::kFormat, api::MakeValue(dst_format)); auto weight_value = dpico::GetTensorInfo(weight_node); MS_ASSERT(weight_value != nullptr); auto weight_shape = weight_value->shape(); @@ -190,7 +198,7 @@ int HandleConstConvWeightShared(const api::FuncGraphPtr &graph, const AnfNodePtr auto abstract = weight_node->abstract(); MS_ASSERT(abstract != nullptr); abstract = abstract->Clone(); - auto shape_ptr = std::make_shared(shape); + auto shape_ptr = api::MakeShared(shape); if (shape_ptr == nullptr) { MS_LOG(ERROR) << "shape ptr is nullptr."; return RET_ERROR; @@ -198,14 +206,15 @@ int HandleConstConvWeightShared(const api::FuncGraphPtr &graph, const AnfNodePtr abstract->set_shape(shape_ptr); trans_cnode->set_abstract(abstract); } - auto post_cnode = post_node->cast(); + auto post_cnode = post_node->cast(); manager->SetEdge(post_cnode, weight_node_user.second, trans_cnode); } return RET_OK; } -int UnifyConstConvWeight(const api::FuncGraphPtr &graph, const AnfNodePtr &weight_node, mindspore::Format src_format, - mindspore::Format dst_format, std::set *has_visited) { +int UnifyConstConvWeight(const api::FuncGraphPtr &graph, const api::AnfNodePtr &weight_node, + mindspore::Format src_format, mindspore::Format dst_format, + std::set *has_visited) { MS_ASSERT(graph != nullptr && weight_node != nullptr && has_visited != nullptr); if (src_format == dst_format) { return lite::RET_OK; @@ -247,15 +256,15 @@ void GetAllFuncGraph(const api::FuncGraphPtr &func_graph, std::setnodes(); for (auto &node : nodes) { - auto new_fg = api::FuncGraph::GetFuncGraphFromAnfNode(node); + auto new_fg = api::GetValueNode(node); if (new_fg != nullptr) { GetAllFuncGraph(new_fg, all_func_graphs); } - if (utils::isa(node)) { - auto cnode = node->cast(); + if (api::utils::isa(node)) { + auto cnode = node->cast(); for (auto &input : cnode->inputs()) { - if (input->isa()) { - new_fg = api::FuncGraph::GetFuncGraphFromAnfNode(node); + if (input->isa()) { + new_fg = api::GetValueNode(node); if (new_fg != nullptr) { GetAllFuncGraph(new_fg, all_func_graphs); } @@ -280,23 +289,23 @@ int PostAdjust(const std::set &all_func_graphs) { return RET_OK; } -int UnifyConvWeightFormat(const api::FuncGraphPtr &graph, const CNodePtr &cnode, mindspore::Format src_format, - mindspore::Format dst_format, std::set *has_visited) { +int UnifyConvWeightFormat(const api::FuncGraphPtr &graph, const api::CNodePtr &cnode, mindspore::Format src_format, + mindspore::Format dst_format, std::set *has_visited) { MS_ASSERT(graph != nullptr && cnode != nullptr && has_visited != nullptr); if (src_format == dst_format) { return lite::RET_OK; } - if (!dpico::CheckPrimitiveType(cnode, prim::kPrimConv2DFusion) && - !dpico::CheckPrimitiveType(cnode, prim::kPrimConv2dTransposeFusion)) { + if (!dpico::CheckPrimitiveType(cnode, api::MakeShared()) && + !dpico::CheckPrimitiveType(cnode, api::MakeShared())) { MS_LOG(ERROR) << "cnode is not a member of convolution's family."; return RET_ERROR; } bool is_const_weight = true; auto weight_node = cnode->input(dpico::kInputIndex2); - if (utils::isa(weight_node)) { + if (api::utils::isa(weight_node)) { is_const_weight = false; - } else if (utils::isa(weight_node)) { - auto weight_param_node = weight_node->cast(); + } else if (api::utils::isa(weight_node)) { + auto weight_param_node = weight_node->cast(); if (!weight_param_node->has_default()) { is_const_weight = false; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.h b/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.h index cd0726cda8..3f36f6a458 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/parser_utils.h @@ -19,10 +19,11 @@ #include #include -#include "ir/anf.h" +#include "mindapi/ir/common.h" +#include "mindapi/ir/anf.h" #include "include/api/format.h" -#include "api/ir/func_graph.h" -#include "utils/log_adapter.h" +#include "mindapi/ir/func_graph.h" +#include "mindapi/base/logging.h" #include "include/errorcode.h" #include "google/protobuf/io/zero_copy_stream_impl.h" #include "google/protobuf/text_format.h" @@ -32,8 +33,8 @@ namespace mindspore { namespace lite { void GetAllFuncGraph(const api::FuncGraphPtr &func_graph, std::set *all_func_graphs); int PostAdjust(const std::set &all_func_graphs); -int UnifyConvWeightFormat(const api::FuncGraphPtr &graph, const CNodePtr &cnode, mindspore::Format src_format, - mindspore::Format dst_format, std::set *has_visited); +int UnifyConvWeightFormat(const api::FuncGraphPtr &graph, const api::CNodePtr &cnode, mindspore::Format src_format, + mindspore::Format dst_format, std::set *has_visited); bool ReadProtoFromCodedInputStream(google::protobuf::io::CodedInputStream *coded_stream, google::protobuf::Message *proto); int ReadProtoFromText(const std::string &file, google::protobuf::Message *message); diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.cc b/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.cc index e52dcea995..6664e0a2b3 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.cc @@ -16,16 +16,22 @@ #include "parser/unify_format.h" #include +#include "common/check_base.h" #include "common/format_utils.h" -#include "ops/op_utils.h" #include "parser/parser_utils.h" +#include "ops/tuple_get_item.h" +#include "ops/adam.h" +#include "ops/sgd.h" +#include "ops/fusion/conv2d_fusion.h" +#include "ops/fusion/conv2d_transpose_fusion.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { -void UnifyFormatToNHWC::GetTransNodeFormatType(const CNodePtr &cnode, dpico::TransTypePair *trans_info) { +void UnifyFormatToNHWC::GetTransNodeFormatType(const api::CNodePtr &cnode, dpico::TransTypePair *trans_info) { MS_ASSERT(cnode != nullptr && trans_info != nullptr); auto prim_node = cnode->input(0); - auto prim = GetValueNode(prim_node); + auto prim = api::GetValueNode(prim_node); MS_ASSERT(prim != nullptr); auto &specify_ops = dpico::GetAssignedFormatOpSet(); if (specify_ops.find(prim->name()) != specify_ops.end()) { @@ -34,10 +40,10 @@ void UnifyFormatToNHWC::GetTransNodeFormatType(const CNodePtr &cnode, dpico::Tra } } -STATUS UnifyFormatToNHWC::GenNewInput(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, std::vector perm, - bool before, size_t index) { +STATUS UnifyFormatToNHWC::GenNewInput(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, + const std::vector &perm, bool before, size_t index) { MS_ASSERT(func_graph != nullptr && cnode != nullptr); - AnfNodePtr trans_input = before ? cnode->input(index) : cnode; + api::AnfNodePtr trans_input = before ? cnode->input(index) : cnode->cast(); std::string trans_name = before ? cnode->fullname_with_scope() + "_pre_" + std::to_string(index - 1) : cnode->fullname_with_scope() + "_post"; auto trans_cnode = dpico::GenTransposeNode(func_graph, trans_input, perm, trans_name); @@ -45,17 +51,17 @@ STATUS UnifyFormatToNHWC::GenNewInput(const api::FuncGraphPtr &func_graph, const if (abstract != nullptr) { trans_cnode->set_abstract(abstract->Clone()); } - auto trans_prim = GetValueNode(trans_cnode->input(0)); + auto trans_prim = api::GetValueNode(trans_cnode->input(0)); if (perm == dpico::kNC2NH) { - trans_prim->AddAttr(ops::kFormat, MakeValue(NCHW)); + trans_prim->AddAttr(ops::kFormat, api::MakeValue(NCHW)); } else if (perm == dpico::kNH2NC) { - trans_prim->AddAttr(ops::kFormat, MakeValue(NHWC)); + trans_prim->AddAttr(ops::kFormat, api::MakeValue(NHWC)); } - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); if (manager == nullptr) { manager = api::FuncGraphManager::Manage(func_graph, true); } - MS_ASSERT(manager != nullptr); + MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "manager is nullptr"); if (before) { manager->SetEdge(cnode, index, trans_cnode); } else { @@ -64,11 +70,11 @@ STATUS UnifyFormatToNHWC::GenNewInput(const api::FuncGraphPtr &func_graph, const return lite::RET_OK; } -STATUS UnifyFormatToNHWC::InsertPreTransNode(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, +STATUS UnifyFormatToNHWC::InsertPreTransNode(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, const std::vector &perm) { MS_ASSERT(func_graph != nullptr && cnode != nullptr); auto prim_node = cnode->input(0); - auto prim = GetValueNode(prim_node); + auto prim = api::GetValueNode(prim_node); MS_ASSERT(prim != nullptr); auto &specify_ops = dpico::GetAssignedFormatOpSet(); if (specify_ops.find(prim->name()) == specify_ops.end()) { @@ -82,28 +88,27 @@ STATUS UnifyFormatToNHWC::InsertPreTransNode(const api::FuncGraphPtr &func_graph return lite::RET_OK; } -STATUS UnifyFormatToNHWC::InsertPostTransNode(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, +STATUS UnifyFormatToNHWC::InsertPostTransNode(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, const std::vector &perm) { MS_ASSERT(func_graph != nullptr && cnode != nullptr); - if (!cnode->abstract()->isa()) { + if (!cnode->abstract()->isa()) { if (GenNewInput(func_graph, cnode, perm, false) != lite::RET_OK) { MS_LOG(ERROR) << "generate a new input failed."; return lite::RET_ERROR; } } else { - MS_ASSERT(func_graph->get_manager() != nullptr); - auto node_map = func_graph->get_manager()->node_users(); - auto &node_users = node_map[cnode]; + MS_CHECK_TRUE_MSG(func_graph->manager() != nullptr, RET_ERROR, "manager is nullptr"); + auto node_users = func_graph->manager()->GetUsers(cnode); for (auto &node_user : node_users) { auto post_node = node_user.first; - if (!dpico::CheckPrimitiveType(post_node, prim::kPrimTupleGetItem)) { + if (!dpico::CheckPrimitiveType(post_node, api::MakeShared())) { MS_LOG(ERROR) << "post node is invalid."; return lite::RET_ERROR; } - if (node_map[post_node].empty()) { + if (func_graph->manager()->GetUsers(post_node).empty()) { continue; } - auto post_cnode = post_node->cast(); + auto post_cnode = post_node->cast(); if (GenNewInput(func_graph, post_cnode, perm, false) != lite::RET_OK) { MS_LOG(ERROR) << "generate a new input failed."; return lite::RET_ERROR; @@ -117,7 +122,7 @@ STATUS UnifyFormatToNHWC::HandleGraphInput(const api::FuncGraphPtr &func_graph) MS_ASSERT(func_graph != nullptr); auto graph_input = func_graph->get_inputs(); for (auto &input : graph_input) { - auto input_param = input->cast(); + auto input_param = input->cast(); MS_ASSERT(input_param != nullptr); auto abstract = input_param->abstract(); MS_ASSERT(abstract != nullptr); @@ -130,29 +135,32 @@ STATUS UnifyFormatToNHWC::HandleGraphInput(const api::FuncGraphPtr &func_graph) continue; } ShapeVector transfer_shape = {shape[0], shape[dpico::kInputIndex2], shape[dpico::kInputIndex3], shape[1]}; - CNodePtr trans_cnode = + api::CNodePtr trans_cnode = dpico::GenTransposeNode(func_graph, input, dpico::kNH2NC, input->fullname_with_scope() + "_nh2nc"); if (trans_cnode == nullptr) { MS_LOG(ERROR) << "create transpose cnode failed."; return lite::RET_ERROR; } - auto trans_prim = GetValueNode(trans_cnode->input(0)); + auto trans_prim = api::GetValueNode(trans_cnode->input(0)); MS_ASSERT(trans_prim != nullptr); - trans_prim->AddAttr(ops::kFormat, MakeValue(NHWC)); + trans_prim->AddAttr(ops::kFormat, api::MakeValue(NHWC)); trans_cnode->set_abstract(abstract->Clone()); - auto transfer_shape_ptr = std::make_shared(transfer_shape); + auto transfer_shape_ptr = api::MakeShared(transfer_shape); if (transfer_shape_ptr == nullptr) { MS_LOG(ERROR) << "transfer_shape_ptr is nullptr."; return RET_ERROR; } abstract->set_shape(transfer_shape_ptr); - MS_ASSERT(func_graph->get_manager() != nullptr); - func_graph->get_manager()->Replace(input, trans_cnode); + MS_CHECK_TRUE_MSG(func_graph->manager() != nullptr, RET_ERROR, "manager is nullptr"); + if (!func_graph->manager()->Replace(input, trans_cnode)) { + MS_LOG(ERROR) << "replace cnode failed."; + return RET_ERROR; + } } return lite::RET_OK; } -STATUS UnifyFormatToNHWC::HandleGraphNode(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode) { +STATUS UnifyFormatToNHWC::HandleGraphNode(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode) { MS_ASSERT(func_graph != nullptr && cnode != nullptr); dpico::TransTypePair trans_info; GetTransNodeFormatType(cnode, &trans_info); @@ -165,15 +173,16 @@ STATUS UnifyFormatToNHWC::HandleGraphNode(const api::FuncGraphPtr &func_graph, c MS_LOG(ERROR) << "insert pre node failed." << cnode->fullname_with_scope(); return lite::RET_ERROR; } - if (dpico::CheckPrimitiveType(cnode, prim::kPrimAdam) || dpico::CheckPrimitiveType(cnode, prim::kPrimSGD)) { + if (dpico::CheckPrimitiveType(cnode, api::MakeShared()) || + dpico::CheckPrimitiveType(cnode, api::MakeShared())) { return lite::RET_OK; } - auto prim = GetValueNode(cnode->input(0)); + auto prim = api::GetValueNode(cnode->input(0)); if (prim == nullptr) { MS_LOG(ERROR) << "current node's prim is nullptr, " << cnode->fullname_with_scope(); return lite::RET_ERROR; } - prim->AddAttr(ops::kFormat, MakeValue(mindspore::NHWC)); + prim->AddAttr(ops::kFormat, api::MakeValue(mindspore::NHWC)); if (InsertPostTransNode(func_graph, cnode, after_perm) != lite::RET_OK) { MS_LOG(ERROR) << "insert post node failed." << cnode->fullname_with_scope(); return lite::RET_ERROR; @@ -191,28 +200,13 @@ bool UnifyFormatToNHWC::BasicProcess(const api::FuncGraphPtr &func_graph, bool m auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); int status; for (auto &node : node_list) { - if (!utils::isa(node)) { + if (!api::utils::isa(node)) { continue; } - auto cnode = node->cast(); + auto cnode = node->cast(); if (dpico::IsSpecialType(cnode)) { continue; } - if (dpico::CheckPrimitiveType(node, prim::kPrimIf) || dpico::CheckPrimitiveType(node, prim::kPrimWhile)) { - auto sub_func_graph = api::FuncGraph::GetFuncGraphFromAnfNode(cnode->input(1)); - if (sub_func_graph == nullptr) { - MS_LOG(ERROR) << "sub graph is nullptr."; - return false; - } - (void)BasicProcess(sub_func_graph, false); - sub_func_graph = api::FuncGraph::GetFuncGraphFromAnfNode(cnode->input(dpico::kInputIndex2)); - if (sub_func_graph == nullptr) { - MS_LOG(ERROR) << "sub graph is nullptr."; - return false; - } - (void)BasicProcess(sub_func_graph, false); - continue; - } status = HandleGraphNode(func_graph, cnode); if (status != lite::RET_OK && status != lite::RET_NO_CHANGE) { return false; @@ -226,37 +220,17 @@ bool UnifyFormatToNHWC::BasicProcess(const api::FuncGraphPtr &func_graph, bool m return true; } -STATUS UnifyFormatToNHWC::ConvWeightFormatTrans(const api::FuncGraphPtr &graph, std::set *has_visited) { +STATUS UnifyFormatToNHWC::ConvWeightFormatTrans(const api::FuncGraphPtr &graph, + std::set *has_visited) { MS_ASSERT(graph != nullptr && has_visited != nullptr); auto node_list = api::FuncGraph::TopoSort(graph->get_return()); for (auto &node : node_list) { - if (!utils::isa(node)) { + if (!api::utils::isa(node)) { continue; } - auto cnode = node->cast(); - if (dpico::CheckPrimitiveType(node, prim::kPrimIf) || dpico::CheckPrimitiveType(node, prim::kPrimWhile)) { - auto sub_func_graph = api::FuncGraph::GetFuncGraphFromAnfNode(cnode->input(1)); - if (sub_func_graph == nullptr) { - MS_LOG(ERROR) << "subgraph is nullptr."; - return false; - } - if (ConvWeightFormatTrans(sub_func_graph, has_visited) != lite::RET_OK) { - MS_LOG(ERROR) << "transform conv weight format failed."; - return lite::RET_ERROR; - } - sub_func_graph = api::FuncGraph::GetFuncGraphFromAnfNode(cnode->input(dpico::kInputIndex2)); - if (sub_func_graph == nullptr) { - MS_LOG(ERROR) << "subgraph is nullptr."; - return false; - } - if (ConvWeightFormatTrans(sub_func_graph, has_visited) != lite::RET_OK) { - MS_LOG(ERROR) << "transform conv weight format failed."; - return lite::RET_ERROR; - } - continue; - } - if (!dpico::CheckPrimitiveType(node, prim::kPrimConv2DFusion) && - !dpico::CheckPrimitiveType(node, prim::kPrimConv2dTransposeFusion)) { + auto cnode = node->cast(); + if (!dpico::CheckPrimitiveType(node, api::MakeShared()) && + !dpico::CheckPrimitiveType(node, api::MakeShared())) { continue; } if (has_visited->find(node) != has_visited->end()) { @@ -276,12 +250,12 @@ bool UnifyFormatToNHWC::Run(const api::FuncGraphPtr &func_graph) { MS_ASSERT(func_graph != nullptr); auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); for (auto &node : node_list) { - auto prim = GetValueNode(node); + auto prim = api::GetValueNode(node); if (prim == nullptr) { continue; } } - std::set has_visited; + std::set has_visited; auto status = ConvWeightFormatTrans(func_graph, &has_visited); if (status != lite::RET_OK) { MS_LOG(ERROR) << "Conv2D weight FormatTrans failed: " << status; diff --git a/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.h b/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.h index 87f6f4cd3e..4bcb2e4262 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.h +++ b/mindspore/lite/tools/converter/adapter/dpico/parser/unify_format.h @@ -35,16 +35,18 @@ class UnifyFormatToNHWC { bool Run(const api::FuncGraphPtr &func_graph); private: - STATUS InsertPostTransNode(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, const std::vector &perm); - STATUS InsertPreTransNode(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, const std::vector &perm); - STATUS GenNewInput(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, std::vector perm, bool before, - size_t index = 0); + STATUS InsertPostTransNode(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, + const std::vector &perm); + STATUS InsertPreTransNode(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, + const std::vector &perm); + STATUS GenNewInput(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, const std::vector &perm, + bool before, size_t index = 0); bool BasicProcess(const api::FuncGraphPtr &func_graph, bool main_graph); - void GetTransNodeFormatType(const CNodePtr &cnode, dpico::TransTypePair *trans_info); + void GetTransNodeFormatType(const api::CNodePtr &cnode, dpico::TransTypePair *trans_info); STATUS HandleGraphInput(const api::FuncGraphPtr &func_graph); - STATUS HandleGraphNode(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode); - STATUS ConvWeightFormatTrans(const api::FuncGraphPtr &graph, std::set *has_visited); - std::map> sub_inputs_map_; + STATUS HandleGraphNode(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode); + STATUS ConvWeightFormatTrans(const api::FuncGraphPtr &graph, std::set *has_visited); + std::map> sub_inputs_map_; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.cc b/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.cc index 52661a5c27..7e507bcd63 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.cc @@ -20,6 +20,7 @@ #include #include #include +#include "ops/tuple_get_item.h" #include "common/anf_util.h" #include "common/string_util.h" #include "common/file_util.h" @@ -59,7 +60,7 @@ int CalibDataGenerator::GenerateDumpConfig(const std::string &dump_cfg_path, return RET_OK; } -std::string CalibDataGenerator::GetInputShapesStr(const AnfNodePtrList &graph_inputs) { +std::string CalibDataGenerator::GetInputShapesStr(const api::AnfNodePtrList &graph_inputs) { std::string input_shapes_str; for (const auto &input : graph_inputs) { ShapeVector shape_vector; @@ -99,7 +100,7 @@ std::string CalibDataGenerator::GetInputShapesStr(const AnfNodePtrList &graph_in return input_shapes_str; } -std::vector CalibDataGenerator::GetInDataFileList(const AnfNodePtrList &graph_inputs) { +std::vector CalibDataGenerator::GetInDataFileList(const api::AnfNodePtrList &graph_inputs) { auto preprocessed_data_dir = DataPreprocessor::GetInstance()->GetPreprocessedDataDir(); if (preprocessed_data_dir.empty()) { MS_LOG(ERROR) << "preprocessed_data_dir is empty."; @@ -136,18 +137,21 @@ int CalibDataGenerator::DumpKernelsData(const std::string &dump_cfg_path, MS_LOG(ERROR) << "set env for dump config failed."; return RET_ERROR; } + std::string benchmark_path = "../../benchmark/benchmark"; + bool use_default_benchmark = true; auto config_info = converter::ConverterContext::GetConfigInfo("dpico"); - if (config_info.empty()) { - MS_LOG(ERROR) << "there is no [dpico] in config file."; - return RET_ERROR; + if (!config_info.empty()) { + if (config_info.find("benchmark_path") != config_info.end()) { + benchmark_path = config_info.at("benchmark_path"); + use_default_benchmark = false; + } } - if (config_info.find("benchmark_path") == config_info.end()) { - MS_LOG(ERROR) << "there is no benchmark_path in [dpico] config section."; - return RET_ERROR; + if (use_default_benchmark) { + MS_LOG(WARNING) << R"(there is no "benchmark_path" in the converter config file, + will use the default value: "../../benchmark/benchmark")"; } - auto benchmark_path = config_info.at("benchmark_path"); - if (benchmark_path.empty()) { - MS_LOG(ERROR) << "benchmark_path content is empty in [dpico] section."; + if (AccessFile(benchmark_path, F_OK) != 0) { + MS_LOG(ERROR) << "File not exist: " << benchmark_path; return RET_ERROR; } std::string current_path; @@ -294,14 +298,14 @@ int CalibDataGenerator::TransBinsToTxt(const std::vector &dump_op_in return RET_OK; } -int CalibDataGenerator::Run(const AnfNodePtrList &graph_inputs, const AnfNodePtrList &nodes) { +int CalibDataGenerator::Run(const api::AnfNodePtrList &graph_inputs, const api::AnfNodePtrList &nodes) { if (graph_inputs.empty()) { MS_LOG(ERROR) << "graph inputs shouldn't be empty."; return RET_ERROR; } auto image_lists = MapperConfigParser::GetInstance()->GetImageLists(); std::vector dump_op_infos; - std::set has_visited; + std::set has_visited; for (const auto &node : nodes) { if (has_visited.find(node) != has_visited.end()) { continue; @@ -313,8 +317,8 @@ int CalibDataGenerator::Run(const AnfNodePtrList &graph_inputs, const AnfNodePtr dump_op_info.input_index = control_flow_inputs_[node].second; } else { dump_op_info.output_index = 0; - if (CheckPrimitiveType(node, prim::kPrimTupleGetItem)) { - auto tuple_get_item_cnode = node->cast(); + if (CheckPrimitiveType(node, api::MakeShared())) { + auto tuple_get_item_cnode = node->cast(); if (tuple_get_item_cnode == nullptr || tuple_get_item_cnode->inputs().size() < kInputIndex2) { MS_LOG(ERROR) << "tuple_get_item_node is invalid. " << node->fullname_with_scope(); return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.h b/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.h index 3ebfd841cc..b5104e4f16 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/calib_data_generator.h @@ -25,9 +25,10 @@ #include #include #include +#include "third_party/securec/include/securec.h" #include "include/registry/converter_context.h" #include "common/data_transpose_utils.h" -#include "ir/anf.h" +#include "mindapi/ir/anf.h" #include "include/errorcode.h" #include "common/check_base.h" @@ -49,15 +50,15 @@ struct DumpOpInfo { class CalibDataGenerator { public: explicit CalibDataGenerator(int dump_level = 0, - const std::map> &control_flow_inputs = {}) + const std::map> &control_flow_inputs = {}) : dump_level_(dump_level), control_flow_inputs_(control_flow_inputs) {} ~CalibDataGenerator() = default; - int Run(const AnfNodePtrList &graph_inputs, const AnfNodePtrList &nodes); + int Run(const api::AnfNodePtrList &graph_inputs, const api::AnfNodePtrList &nodes); private: int GenerateDumpConfig(const std::string &dump_cfg_path, const std::vector &dump_infos); - std::string GetInputShapesStr(const AnfNodePtrList &graph_inputs); - std::vector GetInDataFileList(const AnfNodePtrList &graph_inputs); + std::string GetInputShapesStr(const api::AnfNodePtrList &graph_inputs); + std::vector GetInDataFileList(const api::AnfNodePtrList &graph_inputs); int DumpKernelsData(const std::string &dump_cfg_path, const std::vector &in_data_file_list, const std::string &input_shapes_str); STATUS ParseAttrFromFilename(struct OpAttr *op_attr, const std::string &file_name, bool is_input); @@ -123,7 +124,7 @@ class CalibDataGenerator { return RET_OK; } int dump_level_; - std::map> control_flow_inputs_; + std::map> control_flow_inputs_; }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.cc b/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.cc index bf740a93ae..4f92de5ea6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.cc @@ -20,6 +20,8 @@ #include #include #include +#include "include/api/format.h" +#include "ops/make_tuple.h" #include "common/anf_util.h" #include "common/op_attr.h" #include "common/op_enum.h" @@ -28,6 +30,7 @@ #include "src/om_generator.h" #include "src/graph_split_api.h" #include "ops/tuple_get_item.h" +#include "third_party/securec/include/securec.h" using mindspore::lite::RET_ERROR; using mindspore::lite::RET_OK; @@ -51,16 +54,16 @@ int CheckOmDataCoreInfo(const mapper::DataCoreInfo &data_core_info) { return RET_OK; } -bool CheckInputCNodeSize(const CNodePtr &cnode, size_t *next_idx) { +bool CheckInputCNodeSize(const api::CNodePtr &cnode, size_t *next_idx) { int target_valid_input_size = 1; for (size_t i = 1; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); - if (utils::isa(cnode->input(i))) { - auto param_node = input_node->cast(); + if (api::utils::isa(cnode->input(i))) { + auto param_node = input_node->cast(); if (param_node != nullptr && !param_node->has_default()) { // graph input target_valid_input_size--; } - } else if (utils::isa(input_node)) { + } else if (api::utils::isa(input_node)) { *next_idx = i; target_valid_input_size--; } @@ -68,11 +71,11 @@ bool CheckInputCNodeSize(const CNodePtr &cnode, size_t *next_idx) { return target_valid_input_size == 0; } -bool IsCorrespondOutput(const AnfNodePtr &node, const std::string &target_name) { +bool IsCorrespondOutput(const api::AnfNodePtr &node, const std::string &target_name) { if (node->fullname_with_scope() == target_name) { return true; } - auto cnode = node->cast(); + auto cnode = node->cast(); if (cnode == nullptr) { MS_LOG(INFO) << "cur node isn't a cnode, will stop recursive search. " << node->fullname_with_scope(); return false; @@ -88,8 +91,8 @@ bool IsCorrespondOutput(const AnfNodePtr &node, const std::string &target_name) } } // namespace -CNodePtr CustomOpCreator::CreateCustomOp(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, - const ModelCoreInfoPtr &om_model_info) { +api::CNodePtr CustomOpCreator::CreateCustomOp(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, + const ModelCoreInfoPtr &om_model_info) { MS_CHECK_TRUE_MSG(func_graph != nullptr && subgraph != nullptr && om_model_info != nullptr, nullptr, "obtain nullptr input parameter."); auto om_parameter = CreateOmParameter(func_graph, om_model_info); @@ -99,7 +102,7 @@ CNodePtr CustomOpCreator::CreateCustomOp(const api::FuncGraphPtr &func_graph, Su return nullptr; } - auto prim = std::make_shared(); + auto prim = api::MakeShared(); MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "new Custom failed"); prim->set_type("DPICO"); @@ -127,24 +130,23 @@ CNodePtr CustomOpCreator::CreateCustomOp(const api::FuncGraphPtr &func_graph, Su return custom_cnode; } -ParameterPtr CustomOpCreator::CreateOmParameter(const api::FuncGraphPtr &func_graph, - const ModelCoreInfoPtr &om_model_info) { +api::ParameterPtr CustomOpCreator::CreateOmParameter(const api::FuncGraphPtr &func_graph, + const ModelCoreInfoPtr &om_model_info) { MS_CHECK_TRUE_MSG(func_graph != nullptr && om_model_info != nullptr, nullptr, "obtain nullptr input parameter."); MS_CHECK_TRUE_MSG(om_model_info->modelSize != 0, nullptr, "om model size equals 0."); auto om_parameter = func_graph->add_parameter(); MS_CHECK_TRUE_MSG(om_parameter != nullptr, nullptr, "new parameter failed."); om_parameter->set_name("DPICO_om_data"); - auto type_ptr = TypeIdToType(kNumberTypeUInt8); ShapeVector shape_vector = {static_cast(om_model_info->modelSize)}; - auto abstract_tensor = std::make_shared(type_ptr, shape_vector); + auto abstract_tensor = api::MakeShared(kNumberTypeUInt8, shape_vector); MS_CHECK_TRUE_MSG(abstract_tensor != nullptr, nullptr, "abstract_tensor is nullptr."); om_parameter->set_abstract(abstract_tensor); auto tensor_info = - std::make_shared(kNumberTypeUInt8, ShapeVector({static_cast(om_model_info->modelSize)})); + api::MakeShared(kNumberTypeUInt8, ShapeVector({static_cast(om_model_info->modelSize)})); MS_CHECK_TRUE_MSG(tensor_info != nullptr, nullptr, "tensor_info is nullptr."); - auto tensor_data = tensor_info->data_c(); - MS_CHECK_TRUE_MSG(tensor_data != nullptr, nullptr, "new tensor::Tensor failed."); + auto tensor_data = tensor_info->data(); + MS_CHECK_TRUE_MSG(tensor_data != nullptr, nullptr, "new api::Tensor failed."); MS_CHECK_TRUE_MSG(tensor_info->Size() != 0, nullptr, "tensor size shouldn't be 0"); if (memcpy_s(tensor_data, tensor_info->Size(), om_model_info->modelBuffer, om_model_info->modelSize) != EOK) { MS_LOG(ERROR) << "memcpy_s failed."; @@ -167,7 +169,7 @@ STATUS CustomOpCreator::SetSubgraphInputOutputDims(Subgraph *subgraph, const api } for (const auto &node : subgraph_inputs) { auto node_name = RemoveSpecifiedChar(node->fullname_with_scope(), '\0'); - if (CheckPrimitiveType(node, prim::kPrimCustom)) { + if (CheckPrimitiveType(node, api::MakeShared())) { node_name = GetCustomOutputName(node); MS_CHECK_TRUE_MSG(!node_name.empty(), RET_ERROR, "get custom node output name failed." << node->fullname_with_scope()); @@ -204,7 +206,7 @@ STATUS CustomOpCreator::SetSubgraphInputOutputDims(Subgraph *subgraph, const api } STATUS CustomOpCreator::SetCustomAttrs(const Subgraph &subgraph, const api::FuncGraphPtr &func_graph, - const std::shared_ptr &prim) { + const api::SharedPtr &prim) { MS_CHECK_TRUE_MSG(func_graph != nullptr && prim != nullptr, RET_ERROR, "obtain nullptr input parameter."); std::map> custom_attrs; // add "input_shape " attr @@ -253,9 +255,9 @@ STATUS CustomOpCreator::SetCustomAttrs(const Subgraph &subgraph, const api::Func } STATUS CustomOpCreator::SetCustomOutputs(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, - const CNodePtr &custom_cnode, const ModelCoreInfoPtr &om_model_info) { + const api::CNodePtr &custom_cnode, const ModelCoreInfoPtr &om_model_info) { MS_CHECK_TRUE_MSG(subgraph != nullptr, RET_ERROR, "subgraph is nullptr."); - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "funcgraph manager is nullptr."); auto subgraph_outputs = GetSubgraphOutputs(*subgraph, manager); MS_CHECK_TRUE_MSG(subgraph->outputs_format.size() == subgraph_outputs.size(), RET_ERROR, @@ -277,15 +279,15 @@ STATUS CustomOpCreator::SetCustomOutputs(const api::FuncGraphPtr &func_graph, Su return RET_ERROR; } } - custom_cnode->AddAttr(kOutputsNames, MakeValue(custom_outputs_names)); + custom_cnode->AddAttr(kOutputsNames, api::MakeValue(custom_outputs_names)); return RET_OK; } STATUS CustomOpCreator::SetCustomSingleOutput(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, - const CNodePtr &custom_cnode, + const api::CNodePtr &custom_cnode, const std::shared_ptr &om_model_info, std::vector *custom_outputs_names) { - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "funcgraph manager is nullptr."); auto subgraph_outputs = GetSubgraphOutputs(*subgraph, manager); auto output_info = om_model_info->outputInfos.at(0); @@ -315,16 +317,16 @@ STATUS CustomOpCreator::SetCustomSingleOutput(const api::FuncGraphPtr &func_grap } STATUS CustomOpCreator::SetCustomMultiOutput(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, - const CNodePtr &custom_cnode, + const api::CNodePtr &custom_cnode, const std::shared_ptr &om_model_info, std::vector *custom_outputs_names) { - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "funcgraph manager is nullptr."); auto subgraph_outputs = GetSubgraphOutputs(*subgraph, manager); MS_ASSERT(subgraph->outputs_format.size() == subgraph_outputs.size()); MS_ASSERT(om_model_info->outputInfos.size() >= subgraph_outputs.size()); - AbstractBasePtrList abstract_list; - CNodePtrList subgraph_new_cnodes = {custom_cnode}; + api::AbstractBasePtrList abstract_list; + api::CNodePtrList subgraph_new_cnodes = {custom_cnode}; auto output_formats = subgraph->outputs_format; subgraph->outputs_format.resize(om_model_info->outputInfos.size(), NCHW); size_t has_replace_num = 0; @@ -338,12 +340,12 @@ STATUS CustomOpCreator::SetCustomMultiOutput(const api::FuncGraphPtr &func_graph auto abstract_tensor = CreateTensorAbstract(shape_vector, kDataTypeMap.at(output_info.type)); MS_CHECK_TRUE_MSG(abstract_tensor != nullptr, RET_ERROR, "abstract_tensor is nullptr."); abstract_list.emplace_back(abstract_tensor); - auto tuple_get_item_prim_ptr = std::make_shared(); + auto tuple_get_item_prim_ptr = api::MakeShared(); MS_CHECK_TRUE_MSG(tuple_get_item_prim_ptr != nullptr, RET_ERROR, "tuple_get_item_prim_ptr is nullptr."); auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); - auto get_item_value = NewValueNode(MakeValue(i)); - AnfNodePtrList inputs{tuple_get_item_prim, custom_cnode, get_item_value}; - CNodePtr get_item_cnode = func_graph->NewCNode(inputs); + auto get_item_value = NewValueNode(api::MakeValue(i)); + api::AnfNodePtrList inputs{tuple_get_item_prim, custom_cnode, get_item_value}; + api::CNodePtr get_item_cnode = func_graph->NewCNode(inputs); MS_CHECK_TRUE_MSG(get_item_cnode != nullptr, RET_ERROR, "get_item_cnode is nullptr."); get_item_cnode->set_fullname_with_scope(custom_cnode->fullname_with_scope() + "_getitem_" + std::to_string(i)); auto output_name = RemoveSpecifiedChar(output_info.name, '\0'); @@ -352,7 +354,7 @@ STATUS CustomOpCreator::SetCustomMultiOutput(const api::FuncGraphPtr &func_graph if (has_unsupported_) { auto ori_node_iter = std::find_if( // extra or inconsistent output will be found. subgraph_outputs.begin(), subgraph_outputs.end(), - [output_name](const AnfNodePtr &anf_node) { return IsCorrespondOutput(anf_node, output_name); }); + [output_name](const api::AnfNodePtr &anf_node) { return IsCorrespondOutput(anf_node, output_name); }); if (ori_node_iter == subgraph_outputs.end()) { continue; } @@ -375,7 +377,7 @@ STATUS CustomOpCreator::SetCustomMultiOutput(const api::FuncGraphPtr &func_graph } } else { auto return_cnode = func_graph->get_return(); - if (CheckPrimitiveType(return_cnode->input(1), prim::kPrimMakeTuple)) { + if (CheckPrimitiveType(return_cnode->input(1), api::MakeShared())) { manager->AddEdge(return_cnode->input(1), get_item_cnode); } else { manager->AddEdge(return_cnode, get_item_cnode); @@ -388,7 +390,7 @@ STATUS CustomOpCreator::SetCustomMultiOutput(const api::FuncGraphPtr &func_graph return RET_ERROR; } subgraph->cnodes = subgraph_new_cnodes; - auto abstract_tuple = std::make_shared(abstract_list); + auto abstract_tuple = api::MakeShared(abstract_list); MS_CHECK_TRUE_MSG(abstract_tuple != nullptr, RET_ERROR, "abstract_tuple is nullptr."); custom_cnode->set_abstract(abstract_tuple); return RET_OK; diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.h b/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.h index 9c38e0c654..64d9423a8e 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/custom_creator.h @@ -20,7 +20,7 @@ #include #include #include -#include "api/ir/func_graph.h" +#include "mindapi/ir/func_graph.h" #include "ops/custom.h" #include "src/graph_split_info.h" #include "./op_enum_public.h" @@ -36,21 +36,23 @@ class CustomOpCreator { : custom_id_(custom_id), has_unsupported_(has_unsupported) {} int GetCustomId() { return custom_id_; } ~CustomOpCreator() = default; - CNodePtr CreateCustomOp(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, - const ModelCoreInfoPtr &om_model_info); + api::CNodePtr CreateCustomOp(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, + const ModelCoreInfoPtr &om_model_info); private: - ParameterPtr CreateOmParameter(const api::FuncGraphPtr &func_graph, const ModelCoreInfoPtr &om_model_info); + api::ParameterPtr CreateOmParameter(const api::FuncGraphPtr &func_graph, const ModelCoreInfoPtr &om_model_info); STATUS SetSubgraphInputOutputDims(Subgraph *subgraph, const api::FuncGraphPtr &func_graph, const ModelCoreInfoPtr &om_model_info); STATUS SetCustomAttrs(const Subgraph &subgraph, const api::FuncGraphPtr &func_graph, - const std::shared_ptr &prim); - STATUS SetCustomOutputs(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, const CNodePtr &custom_cnode, + const api::SharedPtr &prim); + STATUS SetCustomOutputs(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, const api::CNodePtr &custom_cnode, const ModelCoreInfoPtr &om_model_info); - STATUS SetCustomSingleOutput(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, const CNodePtr &custom_cnode, - const ModelCoreInfoPtr &om_model_info, std::vector *output_names); - STATUS SetCustomMultiOutput(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, const CNodePtr &custom_cnode, - const ModelCoreInfoPtr &om_model_info, std::vector *output_names); + STATUS SetCustomSingleOutput(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, + const api::CNodePtr &custom_cnode, const ModelCoreInfoPtr &om_model_info, + std::vector *output_names); + STATUS SetCustomMultiOutput(const api::FuncGraphPtr &func_graph, Subgraph *subgraph, + const api::CNodePtr &custom_cnode, const ModelCoreInfoPtr &om_model_info, + std::vector *output_names); int custom_id_; bool has_unsupported_; }; diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.cc b/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.cc index 9f6e9ded73..6b7f26f094 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.cc @@ -284,7 +284,7 @@ int DataPreprocessor::GenerateInputBinFromImages(const std::string &raw_data_pat } return RET_OK; } -int DataPreprocessor::Run(const AnfNodePtrList &inputs) { +int DataPreprocessor::Run(const api::AnfNodePtrList &inputs) { if (inputs.empty()) { MS_LOG(ERROR) << "graph inputs shouldn't be empty."; return RET_ERROR; @@ -298,7 +298,7 @@ int DataPreprocessor::Run(const AnfNodePtrList &inputs) { MS_LOG(ERROR) << "current op don't exist in image_lists. " << op_name; return RET_ERROR; } - auto param_node = input->cast(); + auto param_node = input->cast(); if (param_node == nullptr) { MS_LOG(ERROR) << "graph input node should be parameter ptr. " << input->fullname_with_scope(); return RET_ERROR; diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.h b/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.h index 06b2700277..88c60c963f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/data_preprocessor.h @@ -26,8 +26,8 @@ #include "common/data_transpose_utils.h" #include "include/errorcode.h" #include "common/op_enum.h" -#include "ir/dtype/type_id.h" -#include "ir/anf.h" +#include "mindapi/base/type_id.h" +#include "mindapi/ir/anf.h" #include "src/mapper_config_parser.h" #include "opencv2/core/mat.hpp" #include "common/check_base.h" @@ -39,7 +39,7 @@ namespace dpico { class DataPreprocessor { public: static DataPreprocessor *GetInstance(); - int Run(const AnfNodePtrList &inputs); + int Run(const api::AnfNodePtrList &inputs); const std::string &GetPreprocessedDataDir() const { return preprocessed_data_dir_; } size_t GetBatchSize() const { return batch_size_; } diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/dpico_pass.cc b/mindspore/lite/tools/converter/adapter/dpico/src/dpico_pass.cc index e4fcf7bae2..6229e4701d 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/dpico_pass.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/dpico_pass.cc @@ -22,6 +22,10 @@ #include #include #include +#include "ops/cast.h" +#include "ops/transpose.h" +#include "ops/return.h" +#include "ops/depend.h" #include "common/format_utils.h" #include "common/anf_util.h" #include "common/string_util.h" @@ -39,11 +43,11 @@ namespace mindspore { namespace dpico { namespace { const size_t kMinimumNumbOfSegments = 1; -bool CheckInputDimSize(const CNodePtr &cnode) { +bool CheckInputDimSize(const api::CNodePtr &cnode) { for (size_t i = 0; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); - if (!utils::isa(input_node) && - (input_node->cast() == nullptr || input_node->cast()->has_default())) { + if (!api::utils::isa(input_node) && + (input_node->cast() == nullptr || input_node->cast()->has_default())) { continue; } ShapeVector shape_vector; @@ -58,12 +62,12 @@ bool CheckInputDimSize(const CNodePtr &cnode) { return true; } -bool CheckOpHasInferred(const CNodePtr &cnode) { +bool CheckOpHasInferred(const api::CNodePtr &cnode) { MS_ASSERT(cnode != nullptr); if (GetBoolAttr(cnode, kInferDone)) { return true; } - if (!CheckPrimitiveType(cnode, prim::kPrimTranspose)) { + if (!CheckPrimitiveType(cnode, api::MakeShared())) { return false; } auto abstract = cnode->abstract(); @@ -75,27 +79,27 @@ bool CheckOpHasInferred(const CNodePtr &cnode) { return !shape.empty() && std::all_of(shape.begin(), shape.end(), [](int64_t dim) { return dim > 0; }); } -STATUS AddDumpKernels(const api::FuncGraphPtr &func_graph, const Subgraph &subgraph, AnfNodePtrList *dump_kernels, - std::map> *param_to_cnode) { +STATUS AddDumpKernels(const api::FuncGraphPtr &func_graph, const Subgraph &subgraph, api::AnfNodePtrList *dump_kernels, + std::map> *param_to_cnode) { MS_ASSERT(func_graph != nullptr && dump_kernels != nullptr && param_to_cnode != nullptr); auto subgraph_inputs = GetSubgraphInputs(subgraph, func_graph); MS_CHECK_TRUE_MSG(!subgraph_inputs.empty(), RET_ERROR, "get subgraph inputs failed. subgraph id is " << subgraph.graph_id); dump_kernels->insert(dump_kernels->end(), subgraph_inputs.begin(), subgraph_inputs.end()); bool is_main_graph = - func_graph->get_attr(kIsMainGraph) != nullptr && GetValue(func_graph->get_attr(kIsMainGraph)); + func_graph->get_attr(kIsMainGraph) != nullptr && api::GetValue(func_graph->get_attr(kIsMainGraph)); if (is_main_graph) { return RET_OK; } for (const auto &node : subgraph_inputs) { - if (!utils::isa(node)) { + if (!api::utils::isa(node)) { continue; } if (param_to_cnode->find(node) != param_to_cnode->end()) { continue; } for (const auto &inner_cnode : subgraph.cnodes) { - MS_CHECK_TRUE_MSG(utils::isa(inner_cnode), RET_ERROR, "inner cnode is nullptr"); + MS_CHECK_TRUE_MSG(api::utils::isa(inner_cnode), RET_ERROR, "inner cnode is nullptr"); auto cnode_inputs = inner_cnode->inputs(); auto iter = std::find(cnode_inputs.begin(), cnode_inputs.end(), node); if (iter == cnode_inputs.end()) { @@ -112,11 +116,11 @@ STATUS ModifyGraphInputDataType(const Subgraph &subgraph, const api::FuncGraphPt for (size_t i = 0; i < subgraph_inputs.size(); i++) { auto input = subgraph_inputs.at(i); auto input_node_name = input->fullname_with_scope(); - auto param = input->cast(); + auto param = input->cast(); if (param != nullptr && !param->has_default()) { // only for graph input parameter node auto param_abstract = param->abstract(); MS_CHECK_TRUE_MSG(param_abstract != nullptr, RET_ERROR, "param_abstract is nullptr"); - auto abstractScalar = param_abstract->cast(); + auto abstractScalar = param_abstract->cast(); MS_CHECK_TRUE_MSG(abstractScalar != nullptr, RET_ERROR, "abstractScalar is nullptr"); auto element = abstractScalar->element(); MS_CHECK_TRUE_MSG(element != nullptr, RET_ERROR, "element is nullptr"); @@ -129,19 +133,19 @@ STATUS ModifyGraphInputDataType(const Subgraph &subgraph, const api::FuncGraphPt "can't find \"" << input_node_name << "\" in om model input infos."); switch (correspond_info_iter->type) { case mapper::OpDataType::OP_DTYPE_S8: - element->set_type(kInt8); + element->set_type(api::Type::GetType(kNumberTypeInt8)); break; case mapper::OpDataType::OP_DTYPE_U8: - element->set_type(kUInt8); + element->set_type(api::Type::GetType(kNumberTypeUInt8)); break; case mapper::OpDataType::OP_DTYPE_S16: - element->set_type(kInt16); + element->set_type(api::Type::GetType(kNumberTypeInt16)); break; case mapper::OpDataType::OP_DTYPE_U16: - element->set_type(kUInt16); + element->set_type(api::Type::GetType(kNumberTypeUInt16)); break; case mapper::OpDataType::OP_DTYPE_F32: - element->set_type(kFloat32); + element->set_type(api::Type::GetType(kNumberTypeFloat32)); break; default: MS_LOG(ERROR) << "current op type is unsupported. " << om_model_info->inputInfos.at(i).type; @@ -156,7 +160,7 @@ void PrintUnsupportedOps(const std::map> & size_t unsupported_ops_size, const api::FuncGraphPtr &func_graph) { if (!unsupported_ops.empty()) { if (func_graph->get_attr(kGraphName) != nullptr) { - auto func_graph_name = GetValue(func_graph->get_attr(kGraphName)); + auto func_graph_name = api::GetValue(func_graph->get_attr(kGraphName)); MS_LOG(WARNING) << "func_graph: " << func_graph_name; } MS_LOG(WARNING) << "there are " << unsupported_ops_size << " unsupported ops in this net."; @@ -173,22 +177,33 @@ void PrintUnsupportedOps(const std::map> & #endif } // namespace STATUS DpicoPass::InitDpicoConfigInfo() { + dpico_config_path_ = "./dpico.cfg"; + bool use_default_config = true; auto config_info = converter::ConverterContext::GetConfigInfo("dpico"); - MS_CHECK_TRUE_MSG(!config_info.empty(), RET_ERROR, "there is no [dpico] in config file."); - MS_CHECK_TRUE_MSG(config_info.find("dpico_config_path") != config_info.end(), RET_ERROR, - "there is no dpico_config_path in [dpico] config section."); - dpico_config_path_ = config_info.at("dpico_config_path"); - MS_CHECK_TRUE_MSG(!dpico_config_path_.empty(), RET_ERROR, "dpico_config_path content is empty in [dpico] section."); - if (config_info.find("save_temporary_files") != config_info.end()) { - auto save_temp_file_str = config_info.at("save_temporary_files"); - if (save_temp_file_str == "on") { - save_tmp_files_ = true; - } else if (save_temp_file_str == "off") { - save_tmp_files_ = false; - } else { - MS_LOG(WARNING) << "invalid [save_temporary_files] value, will consider it as off."; - save_tmp_files_ = false; + if (!config_info.empty()) { + if (config_info.find("dpico_config_path") != config_info.end()) { + dpico_config_path_ = config_info.at("dpico_config_path"); + use_default_config = false; } + if (config_info.find("save_temporary_files") != config_info.end()) { + auto save_temp_file_str = config_info.at("save_temporary_files"); + if (save_temp_file_str == "on") { + save_tmp_files_ = true; + } else if (save_temp_file_str == "off") { + save_tmp_files_ = false; + } else { + MS_LOG(WARNING) << "invalid [save_temporary_files] value, will consider it as off."; + save_tmp_files_ = false; + } + } + } + if (use_default_config) { + MS_LOG(WARNING) + << R"(there is no "dpico_config_path" in the converter config file, will use the default value: "./dpico.cfg")"; + } + if (AccessFile(dpico_config_path_, F_OK) != 0) { + MS_LOG(ERROR) << "File not exist: " << dpico_config_path_; + return RET_ERROR; } return RET_OK; } @@ -201,9 +216,9 @@ void DpicoPass::FetchFuncGraphs(const api::FuncGraphPtr &func_graph) { } auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); for (auto &node : node_list) { - auto inner_fg = api::FuncGraph::GetFuncGraphFromAnfNode(node); - if (inner_fg != nullptr) { - FetchFuncGraphs(inner_fg); + auto fg = api::GetValueNode(node); + if (fg != nullptr) { + FetchFuncGraphs(fg); } } } @@ -212,7 +227,7 @@ STATUS DpicoPass::CheckDynamicInputShape(const api::FuncGraphPtr &func_graph) { MS_ASSERT(func_graph != nullptr); auto graph_inputs = func_graph->get_inputs(); for (const auto &node : graph_inputs) { - auto graph_input = node->cast(); + auto graph_input = node->cast(); MS_CHECK_TRUE_MSG(graph_input != nullptr, RET_ERROR, "graph_input is nullptr."); ShapeVector shape_vector; if (GetShapeVectorFromParameter(graph_input, &shape_vector) != RET_OK) { @@ -242,11 +257,11 @@ STATUS DpicoPass::MarkNodes(const api::FuncGraphPtr &func_graph) { std::map> unsupported_ops; size_t unsupported_ops_size = 0; for (auto &node : node_list) { - auto cnode = node->cast(); + auto cnode = node->cast(); if (cnode == nullptr) { continue; } - auto primitive = GetValueNode(cnode->input(0)); + auto primitive = api::GetValueNode(cnode->input(0)); MS_CHECK_TRUE_MSG(primitive != nullptr, RET_ERROR, "primitive is nullptr:" << cnode->fullname_with_scope()); std::string op_type_name; if (GetPrimitiveType(cnode, &op_type_name) != RET_OK) { @@ -258,29 +273,29 @@ STATUS DpicoPass::MarkNodes(const api::FuncGraphPtr &func_graph) { bool is_supported = false; if (IsSpecialType(cnode)) { auto cnode_inputs = cnode->inputs(); - is_supported = CheckPrimitiveType(cnode, prim::kPrimReturn) || CheckPrimitiveType(cnode, prim::kPrimDepend) + is_supported = CheckPrimitiveType(cnode, api::MakeShared()) || + CheckPrimitiveType(cnode, api::MakeShared()) ? false - : std::all_of(cnode_inputs.begin(), cnode_inputs.end(), [](const AnfNodePtr &node) { - return !utils::isa(node) || GetBoolAttr(node, kIsMapperSupported); + : std::all_of(cnode_inputs.begin(), cnode_inputs.end(), [](const api::AnfNodePtr &node) { + return !api::utils::isa(node) || GetBoolAttr(node, kIsMapperSupported); }); } else { auto op_checker = OpCheckerRegistry::GetInstance()->GetOpChecker(op_type_name); if (op_checker != nullptr) { - auto node_map = manager->node_users(); - auto &node_users = node_map[cnode]; + auto node_users = manager->GetUsers(cnode); is_supported = CheckInputDimSize(cnode) && op_checker->Check(cnode, node_users.size(), mindspore::Format::NCHW); } is_supported = is_supported && CheckOpHasInferred(cnode); - is_supported = - is_supported && (CheckPrimitiveType(cnode, prim::kPrimCast) - ? utils::isa(cnode->input(1)) && GetBoolAttr(cnode->input(1), kIsMapperSupported) - : true); + is_supported = is_supported && (CheckPrimitiveType(cnode, api::MakeShared()) + ? api::utils::isa(cnode->input(1)) && + GetBoolAttr(cnode->input(1), kIsMapperSupported) + : true); } - if (!is_supported && !CheckPrimitiveType(cnode, prim::kPrimReturn)) { + if (!is_supported && !CheckPrimitiveType(cnode, api::MakeShared())) { unsupported_ops[op_type_name].push_back(cnode->fullname_with_scope()); unsupported_ops_size++; } - primitive->AddAttr(kIsMapperSupported, MakeValue(is_supported)); + primitive->AddAttr(kIsMapperSupported, api::MakeValue(is_supported)); } #ifdef Debug PrintUnsupportedOps(unsupported_ops, unsupported_ops_size, func_graph); @@ -297,7 +312,7 @@ STATUS DpicoPass::ParseMapperConfig(const api::FuncGraphPtr &func_graph) { std::vector graph_input_names; auto inputs = func_graph->get_inputs(); (void)std::transform(inputs.begin(), inputs.end(), std::back_inserter(graph_input_names), - [](const AnfNodePtr &anode) { return anode->fullname_with_scope(); }); + [](const api::AnfNodePtr &anode) { return anode->fullname_with_scope(); }); if (MapperConfigParser::GetInstance()->Parse(dpico_config_path_, graph_input_names) != RET_OK) { MS_LOG(ERROR) << "parse mapper config file failed."; @@ -320,8 +335,8 @@ STATUS DpicoPass::DataPrepare(const api::FuncGraphPtr &func_graph, bool *use_ori return RET_ERROR; } - AnfNodePtrList dump_kernels; - std::map> param_to_cnode; + api::AnfNodePtrList dump_kernels; + std::map> param_to_cnode; for (auto &graph : func_graphs_) { for (auto &subgraph : graph_split_info_.subgraphs_map[graph]) { if (!subgraph.is_supported) { @@ -333,8 +348,9 @@ STATUS DpicoPass::DataPrepare(const api::FuncGraphPtr &func_graph, bool *use_ori } } } - if (param_to_cnode.empty() && std::all_of(dump_kernels.begin(), dump_kernels.end(), - [](const AnfNodePtr &node) { return utils::isa(node); })) { + if (param_to_cnode.empty() && + std::all_of(dump_kernels.begin(), dump_kernels.end(), + [](const api::AnfNodePtr &node) { return api::utils::isa(node); })) { MS_LOG(DEBUG) << "required tensors are all graph inputs, which do not need to dump data."; return RET_OK; } @@ -430,7 +446,7 @@ STATUS DpicoPass::RemoveTemporaryFiles() { bool DpicoPass::Execute(const api::FuncGraphPtr &func_graph) { MS_CHECK_TRUE_MSG(func_graph != nullptr, false, "func_graph is nullptr."); - func_graph->set_attr(kIsMainGraph, MakeValue(true)); + func_graph->set_attr(kIsMainGraph, api::MakeValue(true)); FetchFuncGraphs(func_graph); auto status = CheckDynamicInputShape(func_graph); if (status == RET_NO_CHANGE) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.cc b/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.cc index 5a05a5b7ca..bf076bf4e6 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.cc @@ -16,19 +16,21 @@ #include "src/dpico_preprocess_pass.h" #include "include/registry/pass_registry.h" -#include "ops/op_utils.h" #include "common/anf_util.h" #include "common/op_enum.h" #include "common/data_transpose_utils.h" #include "ops/fusion/add_fusion.h" +#include "ops/bias_add.h" +#include "ops/op_name.h" +#include "common/check_base.h" #include "common/op_attr.h" namespace mindspore { namespace dpico { namespace { -STATUS InsertTransposeBeforeBiasAdd(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, - const ShapeVector &shape_vector, const abstract::AbstractBasePtr &abstract) { - auto manager = func_graph->get_manager(); +STATUS InsertTransposeBeforeBiasAdd(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, + const ShapeVector &shape_vector, const api::AbstractBasePtr &abstract) { + auto manager = func_graph->manager(); if (manager == nullptr) { MS_LOG(ERROR) << "manager is nullptr. "; return RET_ERROR; @@ -46,40 +48,43 @@ STATUS InsertTransposeBeforeBiasAdd(const api::FuncGraphPtr &func_graph, const C } ShapeVector nc2nh_shape{shape_vector.at(0), shape_vector.at(kInputIndex2), shape_vector.at(kInputIndex3), shape_vector.at(1)}; - auto nc2nh_shape_ptr = std::make_shared(nc2nh_shape); + auto nc2nh_shape_ptr = api::MakeShared(nc2nh_shape); if (nc2nh_shape_ptr == nullptr) { MS_LOG(ERROR) << "new abstract shape failed."; return RET_ERROR; } pre_trans_abstract->set_shape(nc2nh_shape_ptr); pre_trans_cnode->set_abstract(pre_trans_abstract); - auto pre_trans_prim = GetValueNode(pre_trans_cnode->input(0)); + auto pre_trans_prim = api::GetValueNode(pre_trans_cnode->input(0)); MS_ASSERT(pre_trans_prim != nullptr); - pre_trans_prim->AddAttr(ops::kFormat, MakeValue(NCHW)); - pre_trans_prim->AddAttr(kInferDone, MakeValue(true)); + pre_trans_prim->AddAttr(ops::kFormat, api::MakeValue(NCHW)); + pre_trans_prim->AddAttr(kInferDone, api::MakeValue(true)); manager->SetEdge(cnode, kInputIndex1, pre_trans_cnode); return RET_OK; } -STATUS ReplaceBiasAddWithAdd(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, - const PrimitivePtr &primitive) { - auto manager = func_graph->get_manager(); +STATUS ReplaceBiasAddWithAdd(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, + const api::PrimitivePtr &primitive) { + auto manager = func_graph->manager(); if (manager == nullptr) { MS_LOG(ERROR) << "manager is nullptr. "; return RET_ERROR; } - auto prim = std::make_shared(); + auto prim = api::MakeShared(); if (prim == nullptr) { MS_LOG(ERROR) << "new AddFusion failed." << cnode->fullname_with_scope(); return RET_ERROR; } prim->SetAttrs(primitive->attrs()); - prim->AddAttr(ops::kFormat, MakeValue(NHWC)); - auto add_value_node = NewValueNode(prim); + prim->AddAttr(ops::kFormat, api::MakeValue(NHWC)); + auto add_value_node = api::NewValueNode(prim); if (add_value_node == nullptr) { MS_LOG(ERROR) << "new value node failed."; return RET_ERROR; } - (void)manager->Replace(cnode->input(0), add_value_node); + if (!manager->Replace(cnode->input(0), add_value_node)) { + MS_LOG(ERROR) << "replace cnode failed."; + return RET_ERROR; + } auto pre_trans_abstract = GetCNodeInputAbstract(cnode, kInputIndex1); if (pre_trans_abstract == nullptr) { MS_LOG(ERROR) << "cnode input_1 's abstract is nullptr. " << cnode->fullname_with_scope(); @@ -89,9 +94,9 @@ STATUS ReplaceBiasAddWithAdd(const api::FuncGraphPtr &func_graph, const CNodePtr cnode->set_fullname_with_scope(cnode->fullname_with_scope() + "_converted_to_add"); return RET_OK; } -STATUS InsertTransposeAfterBiasAdd(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode, +STATUS InsertTransposeAfterBiasAdd(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode, const ShapeVector &shape_vector) { - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); if (manager == nullptr) { MS_LOG(ERROR) << "manager is nullptr. "; return RET_ERROR; @@ -107,10 +112,10 @@ STATUS InsertTransposeAfterBiasAdd(const api::FuncGraphPtr &func_graph, const CN return RET_ERROR; } post_trans_cnode->set_abstract(post_trans_abstract->Clone()); - auto post_trans_prim = GetValueNode(post_trans_cnode->input(0)); + auto post_trans_prim = api::GetValueNode(post_trans_cnode->input(0)); MS_ASSERT(post_trans_prim != nullptr); - post_trans_prim->AddAttr(ops::kFormat, MakeValue(NHWC)); - post_trans_prim->AddAttr(kInferDone, MakeValue(true)); + post_trans_prim->AddAttr(ops::kFormat, api::MakeValue(NHWC)); + post_trans_prim->AddAttr(kInferDone, api::MakeValue(true)); if (!manager->Replace(cnode, post_trans_cnode)) { MS_LOG(ERROR) << "replace biasadd with add failed." << cnode->fullname_with_scope(); return RET_ERROR; @@ -118,13 +123,13 @@ STATUS InsertTransposeAfterBiasAdd(const api::FuncGraphPtr &func_graph, const CN return RET_OK; } } // namespace -STATUS DpicoPreprocessPass::PreProcessBiadAdd(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode) { - auto manager = func_graph->get_manager(); +STATUS DpicoPreprocessPass::PreProcessBiadAdd(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode) { + auto manager = func_graph->manager(); if (manager == nullptr) { MS_LOG(ERROR) << "manager is nullptr. "; return RET_ERROR; } - auto primitive = GetValueNode(cnode->input(0)); + auto primitive = api::GetValueNode(cnode->input(0)); if (primitive == nullptr) { MS_LOG(ERROR) << "primitive is nullptr:" << cnode->fullname_with_scope(); return RET_ERROR; @@ -177,11 +182,11 @@ bool DpicoPreprocessPass::Execute(const api::FuncGraphPtr &func_graph) { auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); int status; for (const auto &node : node_list) { - auto cnode = node->cast(); + auto cnode = node->cast(); if (cnode == nullptr) { continue; } - if (CheckPrimitiveType(cnode, prim::kPrimBiasAdd)) { + if (CheckPrimitiveType(cnode, api::MakeShared())) { status = PreProcessBiadAdd(func_graph, cnode); if (status != RET_OK && status != RET_NO_CHANGE) { MS_LOG(ERROR) << "preprocess biasadd for dpico failed."; diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.h b/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.h index 881a6d60de..7400215fdc 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/dpico_preprocess_pass.h @@ -38,7 +38,7 @@ class DpicoPreprocessPass : public registry::PassBase { bool Execute(const api::FuncGraphPtr &func_graph) override; private: - STATUS PreProcessBiadAdd(const api::FuncGraphPtr &func_graph, const CNodePtr &cnode); + STATUS PreProcessBiadAdd(const api::FuncGraphPtr &func_graph, const api::CNodePtr &cnode); }; } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.cc b/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.cc index d87fb86a55..2171e7122a 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.cc @@ -18,6 +18,10 @@ #include #include #include +#include "ops/return.h" +#include "ops/transpose.h" +#include "ops/make_tuple.h" +#include "ops/tuple_get_item.h" #include "common/op_attr.h" #include "include/errorcode.h" #include "include/api/format.h" @@ -41,33 +45,34 @@ struct SegmentInfo { size_t right_border; bool is_supported; }; -CNodePtrList GetFuncGraphTotalCNodes(const api::FuncGraphPtr &func_graph) { - CNodePtrList graph_total_cnodes; +api::CNodePtrList GetFuncGraphTotalCNodes(const api::FuncGraphPtr &func_graph) { + api::CNodePtrList graph_total_cnodes; auto node_list = api::FuncGraph::TopoSort(func_graph->get_return()); for (auto &node : node_list) { - auto cnode = node->cast(); + auto cnode = node->cast(); if (cnode != nullptr && - !CheckPrimitiveType(cnode, prim::kPrimTupleGetItem)) { // tuple_get_item may affect split graph + !CheckPrimitiveType(cnode, api::MakeShared())) { // tuple_get_item may affect split graph graph_total_cnodes.push_back(cnode); } } return graph_total_cnodes; } -bool IsTupleGetItemNeeded(const CNodePtr &cnode, const CNodePtr &linked_cnode, const CNodePtrList &total_cnodes) { - return CheckPrimitiveType(cnode, prim::kPrimTupleGetItem) && +bool IsTupleGetItemNeeded(const api::CNodePtr &cnode, const api::CNodePtr &linked_cnode, + const api::CNodePtrList &total_cnodes) { + return CheckPrimitiveType(cnode, api::MakeShared()) && GetBoolAttr(cnode, kIsMapperSupported) == GetBoolAttr(linked_cnode, kIsMapperSupported) && std::find(total_cnodes.begin(), total_cnodes.end(), cnode) == total_cnodes.end(); } -CNodePtrList GetSubgraphTotalCNodes(const api::FuncGraphPtr &func_graph, const CNodePtrList &cnode_list, - const SegmentInfo &segment_info) { - CNodePtrList total_cnodes{}; +api::CNodePtrList GetSubgraphTotalCNodes(const api::FuncGraphPtr &func_graph, const api::CNodePtrList &cnode_list, + const SegmentInfo &segment_info) { + api::CNodePtrList total_cnodes{}; for (size_t i = segment_info.left_border; i <= segment_info.right_border; i++) { auto &cur_cnode = cnode_list[i]; total_cnodes.push_back(cur_cnode); for (const auto &input_node : cur_cnode->inputs()) { - auto input_cnode = input_node->cast(); + auto input_cnode = input_node->cast(); if (input_cnode == nullptr) { continue; } @@ -78,10 +83,9 @@ CNodePtrList GetSubgraphTotalCNodes(const api::FuncGraphPtr &func_graph, const C auto manager = api::FuncGraphManager::Manage(func_graph, true); MS_CHECK_TRUE_MSG(manager != nullptr, {}, "manager is nullptr."); - auto node_map = manager->node_users(); - auto node_users = node_map[cur_cnode]; + auto node_users = manager->GetUsers(cur_cnode); for (const auto &node_user : node_users) { - auto output_cnode = node_user.first->cast(); + auto output_cnode = node_user.first->cast(); if (IsTupleGetItemNeeded(output_cnode, cur_cnode, total_cnodes)) { total_cnodes.push_back(output_cnode); } @@ -90,7 +94,7 @@ CNodePtrList GetSubgraphTotalCNodes(const api::FuncGraphPtr &func_graph, const C return total_cnodes; } -STATUS GetSubgraphNetType(const CNodePtrList &cnodes, OmNetType *om_net_type) { +STATUS GetSubgraphNetType(const api::CNodePtrList &cnodes, OmNetType *om_net_type) { for (const auto &cnode : cnodes) { std::string op_type_name; if (GetPrimitiveType(cnode, &op_type_name) != RET_OK) { @@ -114,12 +118,12 @@ STATUS GetSubgraphNetType(const CNodePtrList &cnodes, OmNetType *om_net_type) { return RET_OK; } -STATUS GenerateSegmentInfos(const CNodePtrList &graph_total_cnodes, std::vector *segment_infos) { +STATUS GenerateSegmentInfos(const api::CNodePtrList &graph_total_cnodes, std::vector *segment_infos) { MS_CHECK_TRUE_MSG(segment_infos != nullptr, RET_ERROR, "segment_infos are nullptr"); size_t start = 0; for (size_t pos = 0; pos < graph_total_cnodes.size(); pos++) { if (pos == graph_total_cnodes.size() - 1) { - if (!CheckPrimitiveType(graph_total_cnodes[pos], prim::kPrimReturn)) { + if (!CheckPrimitiveType(graph_total_cnodes[pos], api::MakeShared())) { MS_LOG(ERROR) << "last cnode should be return node."; return RET_ERROR; } @@ -153,7 +157,8 @@ STATUS ComputeNetworkSegments(const std::vector &segment_infos, Gra return RET_OK; } -std::vector GenerateSubgraphs(const api::FuncGraphPtr &func_graph, const CNodePtrList &graph_total_cnodes, +std::vector GenerateSubgraphs(const api::FuncGraphPtr &func_graph, + const api::CNodePtrList &graph_total_cnodes, const std::vector &segment_infos, size_t *subgraph_cnt) { MS_CHECK_TRUE_MSG(subgraph_cnt != nullptr, {}, "subgraph_cnt is nullptr."); std::vector subgraphs; @@ -171,12 +176,11 @@ std::vector GenerateSubgraphs(const api::FuncGraphPtr &func_graph, con return subgraphs; } -bool FilterMakeTuple(const api::FuncGraphManagerPtr &manager, const Subgraph &subgraph, const CNodePtr &cnode) { - auto node_map = manager->node_users(); - auto node_users = node_map[cnode]; +bool FilterMakeTuple(const api::FuncGraphManagerPtr &manager, const Subgraph &subgraph, const api::CNodePtr &cnode) { + auto node_users = manager->GetUsers(cnode); bool is_subgraph_output = true; for (const auto &node_user : node_users) { - auto output_cnode = node_user.first->cast(); + auto output_cnode = node_user.first->cast(); if (output_cnode == nullptr) { continue; } @@ -187,41 +191,41 @@ bool FilterMakeTuple(const api::FuncGraphManagerPtr &manager, const Subgraph &su return is_subgraph_output; } -bool IsSubgraphParamInput(const AnfNodePtr &front_node, const AnfNodePtrList &subgraph_param_inputs) { - auto param = front_node->cast(); +bool IsSubgraphParamInput(const api::AnfNodePtr &front_node, const api::AnfNodePtrList &subgraph_param_inputs) { + auto param = front_node->cast(); return !param->has_default() && std::find(subgraph_param_inputs.begin(), subgraph_param_inputs.end(), param) == subgraph_param_inputs.end(); } -bool IsSubgraphCNodeInput(const AnfNodePtr &front_node, const Subgraph &subgraph, - const AnfNodePtrList &subgraph_cnode_inputs) { +bool IsSubgraphCNodeInput(const api::AnfNodePtr &front_node, const Subgraph &subgraph, + const api::AnfNodePtrList &subgraph_cnode_inputs) { return std::find(subgraph.cnodes.begin(), subgraph.cnodes.end(), front_node) == subgraph.cnodes.end() && std::find(subgraph_cnode_inputs.begin(), subgraph_cnode_inputs.end(), front_node) == subgraph_cnode_inputs.end(); } -int DetermineOutputFormat(const AnfNodePtr &output_node, Format *format) { - MS_CHECK_TRUE_MSG(output_node != nullptr && output_node->isa(), RET_ERROR, "output node is invalid."); - auto output_cnode = output_node->cast(); +int DetermineOutputFormat(const api::AnfNodePtr &output_node, Format *format) { + MS_CHECK_TRUE_MSG(output_node != nullptr && output_node->isa(), RET_ERROR, "output node is invalid."); + auto output_cnode = output_node->cast(); int64_t local_format = NCHW; auto search_cnode = output_cnode; const int max_search_depth = 10; int loop = 0; // current node may has no format, which can be obtain by transitivity of format. while (loop < max_search_depth) { - auto primitive = GetValueNode(search_cnode->input(0)); + auto primitive = api::GetValueNode(search_cnode->input(0)); if (primitive == nullptr) { break; } if (primitive->GetAttr(kFormat) != nullptr) { - local_format = GetValue(primitive->GetAttr(kFormat)); + local_format = api::GetValue(primitive->GetAttr(kFormat)); break; } auto input_node = search_cnode->input(1); - if (!utils::isa(input_node)) { + if (!api::utils::isa(input_node)) { break; } - search_cnode = input_node->cast(); + search_cnode = input_node->cast(); loop++; } if (local_format < NCHW || local_format > NCW) { @@ -229,7 +233,7 @@ int DetermineOutputFormat(const AnfNodePtr &output_node, Format *format) { return RET_ERROR; } *format = static_cast(local_format); - if (CheckPrimitiveType(output_cnode, prim::kPrimTranspose)) { + if (CheckPrimitiveType(output_cnode, api::MakeShared())) { auto abstract = GetCNodeInputAbstract(output_cnode, 1); MS_CHECK_TRUE_MSG(abstract != nullptr, RET_ERROR, "input's abstract is nullptr."); ShapeVector input_shape; @@ -284,14 +288,14 @@ int GraphSplit(const std::vector &func_graphs, GraphSplitInfo return RET_OK; } -AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphPtr &func_graph) { +api::AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphPtr &func_graph) { MS_CHECK_TRUE_MSG(func_graph != nullptr, {}, "func_graph is nullptr."); - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, {}, "funcgraph manager is nullptr."); - AnfNodePtrList subgraph_param_inputs; - AnfNodePtrList subgraph_cnode_inputs; + api::AnfNodePtrList subgraph_param_inputs; + api::AnfNodePtrList subgraph_cnode_inputs; for (const auto &cnode : subgraph.cnodes) { - if (CheckPrimitiveType(cnode, prim::kPrimMakeTuple)) { + if (CheckPrimitiveType(cnode, api::MakeShared())) { if (FilterMakeTuple(manager, subgraph, cnode)) { continue; } @@ -299,11 +303,11 @@ AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphP for (size_t i = 1; i < cnode->inputs().size(); i++) { auto front_node = cnode->input(i); MS_CHECK_TRUE_MSG(front_node != nullptr, {}, "input node is nullptr."); - if (utils::isa(front_node)) { + if (api::utils::isa(front_node)) { if (IsSubgraphParamInput(front_node, subgraph_param_inputs)) { subgraph_param_inputs.push_back(front_node); } - } else if (utils::isa(front_node)) { + } else if (api::utils::isa(front_node)) { if (IsSubgraphCNodeInput(front_node, subgraph, subgraph_cnode_inputs)) { subgraph_cnode_inputs.push_back(front_node); } @@ -311,7 +315,7 @@ AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphP } } // keep subgraph input as origin graph input order - AnfNodePtrList subgraph_inputs; + api::AnfNodePtrList subgraph_inputs; auto graph_inputs = func_graph->get_inputs(); for (auto &graph_input : graph_inputs) { if (std::find(subgraph_param_inputs.begin(), subgraph_param_inputs.end(), graph_input) == @@ -324,26 +328,25 @@ AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphP return subgraph_inputs; } -AnfNodePtrList GetSubgraphOutputs(const Subgraph &subgraph, const api::FuncGraphManagerPtr &manager) { - AnfNodePtrList subgraph_outputs; - auto node_map = manager->node_users(); +api::AnfNodePtrList GetSubgraphOutputs(const Subgraph &subgraph, const api::FuncGraphManagerPtr &manager) { + api::AnfNodePtrList subgraph_outputs; for (const auto &cnode : subgraph.cnodes) { - auto node_users = node_map[cnode]; + auto node_users = manager->GetUsers(cnode); for (const auto &node_user : node_users) { - auto output_cnode = node_user.first->cast(); + auto output_cnode = node_user.first->cast(); if (output_cnode == nullptr) { continue; } if (std::find(subgraph.cnodes.begin(), subgraph.cnodes.end(), output_cnode) != subgraph.cnodes.end()) { continue; } - if (!CheckPrimitiveType(cnode, prim::kPrimMakeTuple)) { + if (!CheckPrimitiveType(cnode, api::MakeShared())) { subgraph_outputs.push_back(cnode); break; } for (size_t i = 1; i < cnode->inputs().size(); i++) { auto input_node = cnode->input(i); - if (utils::isa(input_node) && + if (api::utils::isa(input_node) && std::find(subgraph.cnodes.begin(), subgraph.cnodes.end(), input_node) != subgraph.cnodes.end() && std::find(subgraph_outputs.begin(), subgraph_outputs.end(), input_node) == subgraph_outputs.end()) { subgraph_outputs.push_back(input_node); @@ -357,7 +360,7 @@ AnfNodePtrList GetSubgraphOutputs(const Subgraph &subgraph, const api::FuncGraph int FillSubgraphOutputsFormat(Subgraph *subgraph, const api::FuncGraphPtr &func_graph) { MS_CHECK_TRUE_MSG(subgraph != nullptr && func_graph != nullptr, RET_ERROR, "output node is invalid."); - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "func_graph's manager is a nullptr."); auto subgraph_outputs = GetSubgraphOutputs(*subgraph, manager); for (auto &output_node : subgraph_outputs) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.h b/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.h index 57ff59f49d..a30acdc482 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_api.h @@ -20,7 +20,7 @@ #include #include #include -#include "api/ir/func_graph_manager.h" +#include "mindapi/ir/common.h" namespace mindspore { namespace dpico { @@ -28,8 +28,8 @@ struct Subgraph; struct GraphSplitInfo; int GraphSplit(const std::vector &func_graphs, GraphSplitInfo *graph_split_info); -AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphPtr &func_graph); -AnfNodePtrList GetSubgraphOutputs(const Subgraph &subgraph, const api::FuncGraphManagerPtr &manager); +api::AnfNodePtrList GetSubgraphInputs(const Subgraph &subgraph, const api::FuncGraphPtr &func_graph); +api::AnfNodePtrList GetSubgraphOutputs(const Subgraph &subgraph, const api::FuncGraphManagerPtr &manager); int FillSubgraphOutputsFormat(Subgraph *subgraph, const api::FuncGraphPtr &func_graph); } // namespace dpico } // namespace mindspore diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_info.h b/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_info.h index cb88c6a351..2cf2d4b1f9 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_info.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/graph_split_info.h @@ -30,11 +30,11 @@ struct Subgraph { int32_t graph_id; bool is_supported; OmNetType om_net_type; - CNodePtrList cnodes; + api::CNodePtrList cnodes; std::vector inputs_dims; std::vector outputs_dims; std::vector outputs_format; - Subgraph(size_t input_id, bool input_flag, OmNetType input_type, CNodePtrList input_cnodes) + Subgraph(size_t input_id, bool input_flag, OmNetType input_type, api::CNodePtrList input_cnodes) : graph_id(input_id), is_supported(input_flag), om_net_type(input_type), cnodes(std::move(input_cnodes)) {} }; diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/mapper_config_parser.cc b/mindspore/lite/tools/converter/adapter/dpico/src/mapper_config_parser.cc index 6f9b1cd00f..1c471d2b87 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/mapper_config_parser.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/mapper_config_parser.cc @@ -20,7 +20,7 @@ #include "common/op_enum.h" #include "common/string_util.h" #include "common/file_util.h" -#include "utils/log_adapter.h" +#include "mindapi/base/logging.h" namespace mindspore { namespace dpico { diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.cc b/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.cc index 1aa8c33d38..e26dc9acd0 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.cc +++ b/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.cc @@ -20,6 +20,9 @@ #include #include #include +#include "ops/tuple_get_item.h" +#include "ops/custom.h" +#include "ops/make_tuple.h" #include "common/anf_util.h" #include "common/string_util.h" #include "mapper/op_mapper_registry.h" @@ -35,29 +38,28 @@ namespace { const std::unordered_map kMapperSupportedTypes = { {kNumberTypeUInt8, "U8"}, {kNumberTypeInt8, "S8"}, {kNumberTypeInt16, "S16"}, {kNumberTypeUInt16, "U16"}, {kNumberTypeFloat16, "FP16"}, {kNumberTypeFloat32, "FP32"}}; -CNodePtrList GetOutputCNodes(const api::FuncGraphManagerPtr &manager, const AnfNodePtr &node) { +api::CNodePtrList GetOutputCNodes(const api::FuncGraphManagerPtr &manager, const api::AnfNodePtr &node) { MS_CHECK_TRUE_MSG(manager != nullptr && node != nullptr, {}, "obtain nullptr input parameter."); - CNodePtrList output_cnodes; - auto node_map = manager->node_users(); - auto &node_users = node_map[node]; + api::CNodePtrList output_cnodes; + auto node_users = manager->GetUsers(node); if (node_users.size() == 1) { - auto output_cnode = node_users.begin()->first->cast(); + auto output_cnode = node_users.begin()->first->cast(); if (output_cnode != nullptr) { output_cnodes.emplace_back(output_cnode); } } else { - std::map output_cnode_ptr_map; + std::map output_cnode_ptr_map; bool has_tuple_get_item = false; for (const auto &node_user : node_users) { - auto output_cnode = node_user.first->cast(); + auto output_cnode = node_user.first->cast(); if (output_cnode != nullptr) { - if (CheckPrimitiveType(output_cnode, prim::kPrimTupleGetItem)) { + if (CheckPrimitiveType(output_cnode, api::MakeShared())) { has_tuple_get_item = true; auto last_input_idx = output_cnode->inputs().size() - 1; auto anode = output_cnode->input(last_input_idx); MS_CHECK_TRUE_MSG(anode != nullptr, {}, output_cnode->fullname_with_scope() << " input(" << last_input_idx << ") is nullptr."); - auto value_node = anode->cast(); + auto value_node = anode->cast(); MS_CHECK_TRUE_MSG(value_node != nullptr, {}, "value node is nullptr. " << anode->fullname_with_scope()); auto value_ptr = value_node->value(); MS_CHECK_TRUE_MSG(value_ptr != nullptr, {}, "value ptr is nullptr. " << anode->fullname_with_scope()); @@ -73,7 +75,7 @@ CNodePtrList GetOutputCNodes(const api::FuncGraphManagerPtr &manager, const AnfN } if (has_tuple_get_item) { (void)std::transform(output_cnode_ptr_map.begin(), output_cnode_ptr_map.end(), std::back_inserter(output_cnodes), - [](const std::pair &iter) { return iter.second; }); + [](const std::pair &iter) { return iter.second; }); } } return output_cnodes; @@ -89,12 +91,11 @@ std::string GetOutNodesStr(const api::FuncGraphManagerPtr &manager, const Subgra out_nodes_str.push_back(';'); } auto subgraph_outputs = GetSubgraphOutputs(sub_graph, manager); - auto node_map = manager->node_users(); - AnfNodePtrList report_nodes; + api::AnfNodePtrList report_nodes; for (const auto &output : subgraph_outputs) { - auto node_users = node_map[output]; + auto node_users = manager->GetUsers(output); for (const auto &node_user : node_users) { - auto output_cnode = node_user.first->cast(); + auto output_cnode = node_user.first->cast(); if (output_cnode == nullptr) { continue; } @@ -104,17 +105,17 @@ std::string GetOutNodesStr(const api::FuncGraphManagerPtr &manager, const Subgra } } out_nodes_str = std::accumulate(report_nodes.begin(), report_nodes.end(), out_nodes_str, - [](const std::string &res, const AnfNodePtr &anf_node_ptr) { + [](const std::string &res, const api::AnfNodePtr &anf_node_ptr) { return res + anf_node_ptr->fullname_with_scope() + ":0;"; }); return out_nodes_str; } -std::string GetInputTypeStr(const AnfNodePtrList &subgraph_inputs, +std::string GetInputTypeStr(const api::AnfNodePtrList &subgraph_inputs, const std::unordered_map &mapper_config) { std::string input_type_str; for (const auto &input : subgraph_inputs) { auto node_name = input->fullname_with_scope(); - if (CheckPrimitiveType(input, prim::kPrimCustom)) { + if (CheckPrimitiveType(input, api::MakeShared())) { node_name = GetCustomOutputName(input); MS_CHECK_TRUE_MSG(!node_name.empty(), {}, "get custom node origin name failed." << input->fullname_with_scope()); } @@ -140,7 +141,7 @@ std::string GetInputTypeStr(const AnfNodePtrList &subgraph_inputs, } return input_type_str; } -STATUS ConfigImageList(const AnfNodePtrList &subgraph_inputs, std::ofstream *mapper_ofs) { +STATUS ConfigImageList(const api::AnfNodePtrList &subgraph_inputs, std::ofstream *mapper_ofs) { MS_CHECK_TRUE_MSG(mapper_ofs != nullptr, RET_ERROR, "mapper_ofs is nullptr."); auto image_lists = MapperConfigParser::GetInstance()->GetImageLists(); if (image_lists.empty()) { @@ -157,7 +158,7 @@ STATUS ConfigImageList(const AnfNodePtrList &subgraph_inputs, std::ofstream *map mapper_ofs->close(); return RET_ERROR; } - if (CheckPrimitiveType(input, prim::kPrimCustom)) { + if (CheckPrimitiveType(input, api::MakeShared())) { node_name = GetCustomOutputName(input); if (node_name.empty()) { MS_LOG(ERROR) << "get custom node origin name failed." << input->fullname_with_scope(); @@ -177,7 +178,7 @@ STATUS ConfigImageList(const AnfNodePtrList &subgraph_inputs, std::ofstream *map } } // namespace -int OmGenerator::GenerateAippConfig(const std::string &aipp_cfg_path, const AnfNodePtrList &subgraph_inputs) { +int OmGenerator::GenerateAippConfig(const std::string &aipp_cfg_path, const api::AnfNodePtrList &subgraph_inputs) { auto aipp_modules = MapperConfigParser::GetInstance()->GetAippModules(); bool need_aipp_cfg = false; std::ofstream aipp_ofs; @@ -218,7 +219,7 @@ int OmGenerator::GenerateAippConfig(const std::string &aipp_cfg_path, const AnfN int OmGenerator::GenerateMapperConfig(const api::FuncGraphPtr &func_graph, const Subgraph &sub_graph, int custom_id, const std::string &mapper_cfg_path) { MS_CHECK_TRUE_MSG(func_graph != nullptr, RET_ERROR, "func_graph is nullptr."); - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "funcgraph manager is nullptr."); auto mapper_config = MapperConfigParser::GetInstance()->GetCommonConfig(); MS_CHECK_TRUE_MSG(!mapper_config.empty(), RET_ERROR, "mapper config is empty."); @@ -271,7 +272,8 @@ int OmGenerator::GenerateMapperConfig(const api::FuncGraphPtr &func_graph, const return RET_OK; } -int OmGenerator::TransformSubGraphInputs(const AnfNodePtrList &inputs, std::vector *base_operators) { +int OmGenerator::TransformSubGraphInputs(const api::AnfNodePtrList &inputs, + std::vector *base_operators) { MS_CHECK_TRUE_MSG(!inputs.empty(), RET_ERROR, "subgraph inputs shouldn't be empty."); MS_CHECK_TRUE_MSG(base_operators != nullptr, RET_ERROR, "base_operators is nullptr."); for (const auto &input : inputs) { @@ -282,7 +284,7 @@ int OmGenerator::TransformSubGraphInputs(const AnfNodePtrList &inputs, std::vect preprocess_operator->SetDimOrderFormat(mapper::DimOrderFormat::NCHW_FORMAT); auto op_name = input->fullname_with_scope(); - if (CheckPrimitiveType(input, prim::kPrimCustom)) { + if (CheckPrimitiveType(input, api::MakeShared())) { op_name = GetCustomOutputName(input); MS_CHECK_TRUE_MSG(!op_name.empty(), RET_ERROR, "get custom node output name failed." << input->fullname_with_scope()); @@ -303,16 +305,17 @@ int OmGenerator::TransformSubGraphInputs(const AnfNodePtrList &inputs, std::vect return RET_OK; } -int OmGenerator::TransformSubGraphCNodes(const api::FuncGraphManagerPtr &manager, const CNodePtrList &cnodes, +int OmGenerator::TransformSubGraphCNodes(const api::FuncGraphManagerPtr &manager, const api::CNodePtrList &cnodes, std::vector *base_operators) { MS_CHECK_TRUE_MSG(!cnodes.empty(), RET_ERROR, "subgraph inputs shouldn't be empty."); MS_CHECK_TRUE_MSG(base_operators != nullptr, RET_ERROR, "base_operators is nullptr."); for (const auto &cnode : cnodes) { - MS_CHECK_TRUE_MSG(utils::isa(cnode), RET_ERROR, "cur node should be a cnode"); - auto primitive = GetValueNode(cnode->input(0)); + MS_CHECK_TRUE_MSG(api::utils::isa(cnode), RET_ERROR, "cur node should be a cnode"); + auto primitive = api::GetValueNode(cnode->input(0)); MS_CHECK_TRUE_MSG(primitive != nullptr, RET_ERROR, "invalid anf node, which don't have primitive. " << cnode->fullname_with_scope()); - if (CheckPrimitiveType(cnode, prim::kPrimMakeTuple) || CheckPrimitiveType(cnode, prim::kPrimTupleGetItem)) { + if (CheckPrimitiveType(cnode, api::MakeShared()) || + CheckPrimitiveType(cnode, api::MakeShared())) { MS_LOG(DEBUG) << "MakeTuple and TupleGetItem don't need to transform."; continue; } @@ -338,7 +341,7 @@ int OmGenerator::TransformSubGraphCNodes(const api::FuncGraphManagerPtr &manager int OmGenerator::Run(const api::FuncGraphPtr &func_graph, const Subgraph &sub_graph, int custom_id, mapper::ModelCoreInfo *om_model_info, bool use_origin_config) { MS_CHECK_TRUE_MSG(func_graph != nullptr && om_model_info != nullptr, RET_ERROR, "obtain nullptr input parameter."); - auto manager = func_graph->get_manager(); + auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(manager != nullptr, RET_ERROR, "funcgraph manager is nullptr."); std::string mapper_cfg_path; if (!use_origin_config) { diff --git a/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.h b/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.h index 4ed28f62b1..9a07fc022f 100644 --- a/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.h +++ b/mindspore/lite/tools/converter/adapter/dpico/src/om_generator.h @@ -22,7 +22,7 @@ #include #include #include -#include "api/ir/func_graph_manager.h" +#include "mindapi/ir/common.h" #include "src/graph_split_info.h" #include "mapper/op_mapper.h" #include "op/base_operator.h" @@ -37,11 +37,11 @@ class OmGenerator { mapper::ModelCoreInfo *om_model_info, bool use_origin_config); private: - int GenerateAippConfig(const std::string &aipp_cfg_path, const AnfNodePtrList &subgraph_inputs); + int GenerateAippConfig(const std::string &aipp_cfg_path, const api::AnfNodePtrList &subgraph_inputs); int GenerateMapperConfig(const api::FuncGraphPtr &func_graph, const Subgraph &sub_graph, int custom_id, const std::string &cfg); - int TransformSubGraphInputs(const AnfNodePtrList &nodes, std::vector *base_operators); - int TransformSubGraphCNodes(const api::FuncGraphManagerPtr &manager, const CNodePtrList &Cnodes, + int TransformSubGraphInputs(const api::AnfNodePtrList &nodes, std::vector *base_operators); + int TransformSubGraphCNodes(const api::FuncGraphManagerPtr &manager, const api::CNodePtrList &Cnodes, std::vector *base_operators); }; } // namespace dpico diff --git a/mindspore/lite/tools/converter/anf_transform.cc b/mindspore/lite/tools/converter/anf_transform.cc index c11a11961d..2fdbee903b 100644 --- a/mindspore/lite/tools/converter/anf_transform.cc +++ b/mindspore/lite/tools/converter/anf_transform.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/anf_transform.h" #include #include diff --git a/mindspore/lite/tools/converter/converter.cc b/mindspore/lite/tools/converter/converter.cc index 0f64f93878..7f1dea9653 100644 --- a/mindspore/lite/tools/converter/converter.cc +++ b/mindspore/lite/tools/converter/converter.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/converter.h" #include #include @@ -48,6 +49,10 @@ void InitConverterParameters(const converter::Flags &flag, converter::ConverterP converter_parameters->model_file = flag.modelFile; converter_parameters->weight_file = flag.weightFile; } +FuncGraphPtr ConvertGraph(const api::FuncGraphPtr &func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} } // namespace FuncGraphPtr Converter::BuildFuncGraph(const converter::Flags &flag) { @@ -57,7 +62,7 @@ FuncGraphPtr Converter::BuildFuncGraph(const converter::Flags &flag) { kernel::PopulateTrainParameters(); #endif MindsporeImporter ms_import; - func_graph_base = ms_import.ImportMindIR(flag); + func_graph_base = api::MakeShared(ms_import.ImportMindIR(flag)); } else { model_parser_ = registry::ModelParserRegistry::GetModelParser(flag.fmk); if (model_parser_ == nullptr) { @@ -72,7 +77,7 @@ FuncGraphPtr Converter::BuildFuncGraph(const converter::Flags &flag) { ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_NOT_SUPPORT); return nullptr; } - auto func_graph = std::dynamic_pointer_cast(func_graph_base); + auto func_graph = ConvertGraph(func_graph_base); if (func_graph == nullptr) { MS_LOG(ERROR) << "func graph is invalid."; return nullptr; diff --git a/mindspore/lite/tools/converter/converter.h b/mindspore/lite/tools/converter/converter.h index 45da07c873..76060b29e1 100644 --- a/mindspore/lite/tools/converter/converter.h +++ b/mindspore/lite/tools/converter/converter.h @@ -17,6 +17,7 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_CONVERTER_H #define MINDSPORE_LITE_TOOLS_CONVERTER_CONVERTER_H +#define USE_DEPRECATED_API #include #include #include "include/registry/model_parser.h" diff --git a/mindspore/lite/tools/converter/export_model.cc b/mindspore/lite/tools/converter/export_model.cc index 28138679ec..128506522a 100644 --- a/mindspore/lite/tools/converter/export_model.cc +++ b/mindspore/lite/tools/converter/export_model.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/export_model.h" #include #include diff --git a/mindspore/lite/tools/converter/import/mindir_control_flow_adjust.cc b/mindspore/lite/tools/converter/import/mindir_control_flow_adjust.cc index c05ff9c75f..a5b24eb0d0 100644 --- a/mindspore/lite/tools/converter/import/mindir_control_flow_adjust.cc +++ b/mindspore/lite/tools/converter/import/mindir_control_flow_adjust.cc @@ -136,7 +136,9 @@ FuncGraphPtr MindIRControlFlowAdjust::AddAfterFuncGraph(const FuncGraphPtr &fg, MS_LOG(ERROR) << "new MakeTuple failed"; return nullptr; } - auto make_tuple_prim = NewValueNode(make_tuple_prim_ptr); + auto make_tuple_prim_c = make_tuple_prim_ptr->GetPrim(); + MS_CHECK_TRUE_MSG(make_tuple_prim_c != nullptr, nullptr, "Failed to create make_tuple_prim_c."); + auto make_tuple_prim = NewValueNode(make_tuple_prim_c); MS_CHECK_TRUE_MSG(make_tuple_prim != nullptr, nullptr, "Failed to create value node."); make_tuple_inputs.insert(make_tuple_inputs.begin(), make_tuple_prim); auto make_tuple_cnode = after_fg->NewCNode(make_tuple_inputs); @@ -147,7 +149,9 @@ FuncGraphPtr MindIRControlFlowAdjust::AddAfterFuncGraph(const FuncGraphPtr &fg, MS_LOG(ERROR) << "new Return failed"; return nullptr; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + MS_CHECK_TRUE_MSG(return_prim_c != nullptr, nullptr, "Failed to create return_prim_c."); + auto value_node = NewValueNode(return_prim_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "Failed to create value node."); std::vector op_inputs = {value_node, make_tuple_cnode}; auto cnode = after_fg->NewCNode(op_inputs); @@ -160,7 +164,9 @@ FuncGraphPtr MindIRControlFlowAdjust::AddAfterFuncGraph(const FuncGraphPtr &fg, MS_LOG(ERROR) << "new Return failed"; return nullptr; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + MS_CHECK_TRUE_MSG(return_prim_c != nullptr, nullptr, "Failed to create return_prim_c."); + auto value_node = NewValueNode(return_prim_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "Failed to create value node."); std::vector op_inputs{value_node, after_fg->get_inputs().front()}; auto return_cnode = after_fg->NewCNode(op_inputs); diff --git a/mindspore/lite/tools/converter/import/mindspore_importer.cc b/mindspore/lite/tools/converter/import/mindspore_importer.cc index f777b4d1bd..73bddace67 100644 --- a/mindspore/lite/tools/converter/import/mindspore_importer.cc +++ b/mindspore/lite/tools/converter/import/mindspore_importer.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/import/mindspore_importer.h" #include #include diff --git a/mindspore/lite/tools/converter/import/primitive_adjust.cc b/mindspore/lite/tools/converter/import/primitive_adjust.cc index b52a4ac071..f6f34c4145 100644 --- a/mindspore/lite/tools/converter/import/primitive_adjust.cc +++ b/mindspore/lite/tools/converter/import/primitive_adjust.cc @@ -14,11 +14,13 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/import/primitive_adjust.h" #include #include #include #include +#include "ops/op_utils.h" #include "ops/batch_norm.h" #include "ops/elu.h" #include "ops/fused_batch_norm.h" @@ -226,7 +228,8 @@ int MoveAttrMapCommon(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + T dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); value_node->set_value(dst_prim); @@ -241,7 +244,8 @@ int MoveAttrMapActivation(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto act_prim = std::make_shared(); + ops::Activation act_node; + auto act_prim = act_node.GetPrim(); MS_CHECK_TRUE_MSG(act_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); act_prim->SetAttrs(src_prim->attrs()); auto iter = activation_map.find(src_prim->name()); @@ -249,7 +253,7 @@ int MoveAttrMapActivation(const CNodePtr &cnode) { MS_LOG(ERROR) << "activation mode is unsupported."; return lite::RET_ERROR; } - act_prim->set_activation_type(iter->second); + act_node.set_activation_type(iter->second); value_node->set_value(act_prim); return lite::RET_OK; } @@ -262,7 +266,8 @@ int MoveAttrMapActivationGrad(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto act_grad_prim = std::make_shared(); + ops::ActivationGrad act_grad_node; + auto act_grad_prim = act_grad_node.GetPrim(); MS_CHECK_TRUE_MSG(act_grad_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); act_grad_prim->SetAttrs(src_prim->attrs()); auto iter = activation_map.find(src_prim->name()); @@ -270,7 +275,7 @@ int MoveAttrMapActivationGrad(const CNodePtr &cnode) { MS_LOG(ERROR) << "activation mode is unsupported."; return lite::RET_ERROR; } - act_grad_prim->set_activation_type(iter->second); + act_grad_node.set_activation_type(iter->second); value_node->set_value(act_grad_prim); return lite::RET_OK; } @@ -284,7 +289,8 @@ int MoveAttrMapReduce(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + ops::ReduceFusion dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); auto iter = reduce_map.find(src_prim->name()); @@ -292,8 +298,8 @@ int MoveAttrMapReduce(const CNodePtr &cnode) { MS_LOG(ERROR) << "reduce mode is unsupported."; return lite::RET_ERROR; } - dst_prim->set_mode(iter->second); - dst_prim->set_coeff(1.0f); + dst_node.set_mode(iter->second); + dst_node.set_coeff(1.0f); value_node->set_value(dst_prim); return lite::RET_OK; } @@ -330,9 +336,11 @@ int MoveAttrMapConv2D(const CNodePtr &cnode) { } PrimitivePtr dst_prim{nullptr}; if (opt::CheckPrimitiveType(cnode, prim::kPrimConv2D)) { - dst_prim = std::make_shared(); + ops::Conv2DFusion node; + dst_prim = node.GetPrim(); } else if (opt::CheckPrimitiveType(cnode, prim::kPrimConv2DTranspose)) { - dst_prim = std::make_shared(); + ops::Conv2dTransposeFusion node; + dst_prim = node.GetPrim(); } MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); @@ -368,10 +376,12 @@ int MoveAttrPool(const CNodePtr &cnode) { } PrimitivePtr dst_prim = nullptr; if (src_prim->name() == kNameAvgPool) { - dst_prim = std::make_shared(); + ops::AvgPoolFusion dst_node; + dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); } else if (src_prim->name() == kNameMaxPool) { - dst_prim = std::make_shared(); + ops::MaxPoolFusion dst_node; + dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); } else { MS_LOG(ERROR) << "unsupported pooling type."; @@ -405,10 +415,12 @@ int MoveAttrPoolGrad(const CNodePtr &cnode) { PrimitivePtr dst_prim = nullptr; if (src_prim->name() == kNameAvgPoolGrad || src_prim->name() == kNameAvgPoolGradGpu || src_prim->name() == kNameAvgPoolGradCpu) { - dst_prim = std::make_shared(); + ops::AvgPoolGrad dst_node; + dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); } else if (src_prim->name() == kNameMaxPoolGrad) { - dst_prim = std::make_shared(); + ops::MaxPoolGrad dst_node; + dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); } else { MS_LOG(ERROR) << "unsupported pooling type."; @@ -438,7 +450,8 @@ int MoveAttrMapAdder(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + ops::AdderFusion dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); auto status = AdjustConvAttr(dst_prim); @@ -459,12 +472,13 @@ int MoveAttrMapLayerNorm(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + ops::LayerNormFusion dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); dst_prim->SetAttrs(src_prim->attrs()); - dst_prim->set_elementwise_affine(true); + dst_node.set_elementwise_affine(true); if (dst_prim->GetAttr(ops::kEpsilon) == nullptr) { - dst_prim->set_epsilon(1e-7); + dst_node.set_epsilon(1e-7); } value_node->set_value(dst_prim); return lite::RET_OK; @@ -479,19 +493,20 @@ int MoveAttrMapResize(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + ops::Resize dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); auto size = GetValue>(src_prim->GetAttr(ops::kSize)); MS_CHECK_TRUE_MSG(size.size() > 1, RET_ERROR, "out of range."); - dst_prim->set_new_height(size[0]); - dst_prim->set_new_width(size[1]); + dst_node.set_new_height(size[0]); + dst_node.set_new_width(size[1]); if (src_prim->GetAttr(ops::kAlignCorners) != nullptr && GetValue(src_prim->GetAttr(ops::kAlignCorners))) { - dst_prim->set_coordinate_transform_mode(mindspore::ALIGN_CORNERS); + dst_node.set_coordinate_transform_mode(mindspore::ALIGN_CORNERS); } if (src_prim->name() == kNameResizeBilinear) { - dst_prim->set_method(ResizeMethod::LINEAR); + dst_node.set_method(ResizeMethod::LINEAR); } else if (src_prim->name() == kNameResizeNearestNeighbor) { - dst_prim->set_method(ResizeMethod::NEAREST); + dst_node.set_method(ResizeMethod::NEAREST); } value_node->set_value(dst_prim); return lite::RET_OK; @@ -506,7 +521,8 @@ int MoveAttrSlice(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + ops::SliceFusion dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); auto begin = GetValueNode(cnode->input(opt::kInputIndexTwo)); auto begin_value = GetValue>(begin); @@ -515,7 +531,7 @@ int MoveAttrSlice(const CNodePtr &cnode) { for (size_t i = 0; i < begin_value.size(); i++) { axes[i] = static_cast(i); } - dst_prim->set_axes(axes); + dst_node.set_axes(axes); dst_prim->SetAttrs(src_prim->attrs()); value_node->set_value(dst_prim); return lite::RET_OK; @@ -530,13 +546,14 @@ int MoveAttrMapResizeGrad(const CNodePtr &cnode) { MS_LOG(ERROR) << "value node is invalid."; return lite::RET_ERROR; } - auto dst_prim = std::make_shared(); + ops::ResizeGrad dst_node; + auto dst_prim = dst_node.GetPrim(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); if (src_prim->name() == kNameResizeBilinearGrad) { - dst_prim->set_method(ResizeMethod::LINEAR); + dst_node.set_method(ResizeMethod::LINEAR); } else if (src_prim->name() == kNameResizeNearestNeighborGrad) { - dst_prim->set_method(ResizeMethod::NEAREST); + dst_node.set_method(ResizeMethod::NEAREST); } else { MS_LOG(ERROR) << "Resize grad method " << src_prim->name() << "is not supported"; return lite::RET_ERROR; @@ -544,7 +561,7 @@ int MoveAttrMapResizeGrad(const CNodePtr &cnode) { MS_CHECK_TRUE_MSG(src_prim->GetAttr(ops::kAlignCorners) != nullptr, RET_NULL_PTR, "src_prim->GetAttr(ops::kAlignCorners) is nullptr."); auto align_corners = GetValue(src_prim->GetAttr(ops::kAlignCorners)); - dst_prim->set_align_corners(align_corners); + dst_node.set_align_corners(align_corners); value_node->set_value(dst_prim); return lite::RET_OK; } @@ -560,10 +577,12 @@ int MoveAttrBatchNorm(const CNodePtr &cnode) { } auto dst_prim = std::make_shared(); MS_CHECK_TRUE_MSG(dst_prim != nullptr, RET_NULL_PTR, "dst_prim is nullptr."); - dst_prim->SetAttrs(src_prim->attrs()); + auto dst_prim_c = dst_prim->GetPrim(); + MS_CHECK_TRUE_MSG(dst_prim_c != nullptr, RET_NULL_PTR, "dst_prim_c is nullptr."); + dst_prim_c->SetAttrs(src_prim->attrs()); bool is_training = GetValue(src_prim->GetAttr(ops::kIsTraining)); dst_prim->set_mode(static_cast(is_training)); - value_node->set_value(dst_prim); + value_node->set_value(dst_prim_c); return lite::RET_OK; } } // namespace diff --git a/mindspore/lite/tools/converter/legacy_optimizer/graph/convert_fp32_to_fp16_pass.cc b/mindspore/lite/tools/converter/legacy_optimizer/graph/convert_fp32_to_fp16_pass.cc index b832d0cc43..0a4267e842 100644 --- a/mindspore/lite/tools/converter/legacy_optimizer/graph/convert_fp32_to_fp16_pass.cc +++ b/mindspore/lite/tools/converter/legacy_optimizer/graph/convert_fp32_to_fp16_pass.cc @@ -30,7 +30,7 @@ namespace mindspore { namespace lite { namespace { constexpr int kFp16ToFp32Multiply = 2; -} +} // namespace STATUS ConvertFP32ToFP16Pass::Run(schema::MetaGraphT *graph) { if (!need_convert_) { diff --git a/mindspore/lite/tools/converter/legacy_optimizer/graph/dtype_trans_pass.cc b/mindspore/lite/tools/converter/legacy_optimizer/graph/dtype_trans_pass.cc index e05b8979ab..a6e846542f 100644 --- a/mindspore/lite/tools/converter/legacy_optimizer/graph/dtype_trans_pass.cc +++ b/mindspore/lite/tools/converter/legacy_optimizer/graph/dtype_trans_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/legacy_optimizer/graph/dtype_trans_pass.h" #include #include diff --git a/mindspore/lite/tools/converter/legacy_optimizer/graph/infer_quant_param_pass.cc b/mindspore/lite/tools/converter/legacy_optimizer/graph/infer_quant_param_pass.cc index 92e82b18d5..81b751e793 100644 --- a/mindspore/lite/tools/converter/legacy_optimizer/graph/infer_quant_param_pass.cc +++ b/mindspore/lite/tools/converter/legacy_optimizer/graph/infer_quant_param_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/legacy_optimizer/graph/infer_quant_param_pass.h" #include #include diff --git a/mindspore/lite/tools/converter/legacy_optimizer/graph/infershape_pass.cc b/mindspore/lite/tools/converter/legacy_optimizer/graph/infershape_pass.cc index 5b3f7e4d90..3431b18755 100644 --- a/mindspore/lite/tools/converter/legacy_optimizer/graph/infershape_pass.cc +++ b/mindspore/lite/tools/converter/legacy_optimizer/graph/infershape_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/legacy_optimizer/graph/infershape_pass.h" #include #include @@ -79,7 +80,7 @@ void FreeTensors(std::vector *input_tensors, std::vector *ou namespace { constexpr int kBytesPerInt = 4; -} +} // namespace void ConvertTensorList(const MetaGraphT *graph, uint32_t index, bool *convert_succ, std::vector *lite_tensors) { diff --git a/mindspore/lite/tools/converter/legacy_optimizer/graph/tensor_quant_pass.cc b/mindspore/lite/tools/converter/legacy_optimizer/graph/tensor_quant_pass.cc index 01b3b9f8f8..55b9ce86e4 100644 --- a/mindspore/lite/tools/converter/legacy_optimizer/graph/tensor_quant_pass.cc +++ b/mindspore/lite/tools/converter/legacy_optimizer/graph/tensor_quant_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/legacy_optimizer/graph/tensor_quant_pass.h" #include #include diff --git a/mindspore/lite/tools/converter/legacy_optimizer/graph/topological_sort_pass.cc b/mindspore/lite/tools/converter/legacy_optimizer/graph/topological_sort_pass.cc index 9b70fa9b96..dd514bde62 100644 --- a/mindspore/lite/tools/converter/legacy_optimizer/graph/topological_sort_pass.cc +++ b/mindspore/lite/tools/converter/legacy_optimizer/graph/topological_sort_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include #include diff --git a/mindspore/lite/tools/converter/main.cc b/mindspore/lite/tools/converter/main.cc index 26a2247d46..0531f1eee0 100644 --- a/mindspore/lite/tools/converter/main.cc +++ b/mindspore/lite/tools/converter/main.cc @@ -17,6 +17,7 @@ #if defined(__linux__) && !defined(Debug) #include #endif +#define USE_DEPRECATED_API #include "tools/converter/converter.h" #if defined(__linux__) && !defined(Debug) diff --git a/mindspore/lite/tools/converter/micro/coder/opcoders/base/conv2d_base_coder.cc b/mindspore/lite/tools/converter/micro/coder/opcoders/base/conv2d_base_coder.cc index 1615b10060..5e6c17f4fd 100644 --- a/mindspore/lite/tools/converter/micro/coder/opcoders/base/conv2d_base_coder.cc +++ b/mindspore/lite/tools/converter/micro/coder/opcoders/base/conv2d_base_coder.cc @@ -24,7 +24,7 @@ namespace mindspore::lite::micro { namespace { constexpr int kRoundUp = 2; -} +} // namespace Conv2DBaseCoder::~Conv2DBaseCoder() { FreeConvQuantParams(); conv_param_ = nullptr; diff --git a/mindspore/lite/tools/converter/optimizer_manager.cc b/mindspore/lite/tools/converter/optimizer_manager.cc index 6b4a4388f2..9d49396384 100644 --- a/mindspore/lite/tools/converter/optimizer_manager.cc +++ b/mindspore/lite/tools/converter/optimizer_manager.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/optimizer_manager.h" #include #include @@ -47,7 +48,8 @@ bool RunOptimizerPass(const FuncGraphPtr &func_graph, const std::vectorExecute(func_graph)) { + auto api_graph = api::MakeShared(func_graph); + if (!pass_outer->Execute(api_graph)) { MS_LOG(WARNING) << "run pass failed, pass name is " << pass_name; return false; } diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.cc index 9db8ecf39c..aff71d61f4 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::RELU); @@ -39,36 +39,36 @@ ops::PrimitiveC *CaffeReluParser::Parse(const caffe::LayerParameter &proto, cons } } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *CaffeRelu6Parser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeRelu6Parser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::RELU6); prim->set_min_val(0); prim->set_max_val(kValueThreshold6); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *CaffeSigmoidParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeSigmoidParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::SIGMOID); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *CaffeTanhParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeTanhParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::TANH); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *CaffeEluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeEluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::ELU); @@ -80,7 +80,7 @@ ops::PrimitiveC *CaffeEluParser::Parse(const caffe::LayerParameter &proto, const } } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeReluParser("ReLU", new CaffeReluParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.h index ca75301dd2..44db947e9a 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_activation_parser.h @@ -16,6 +16,7 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_PARSER_CAFFE_CAFFE_ACTIVATION_PARSER_H_ #define MINDSPORE_LITE_TOOLS_CONVERTER_PARSER_CAFFE_CAFFE_ACTIVATION_PARSER_H_ +#define USE_DEPRECATED_API #include #include "tools/converter/parser/caffe/caffe_node_parser.h" @@ -28,7 +29,7 @@ class CaffeReluParser : public CaffeNodeParser { CaffeReluParser() : CaffeNodeParser("relu") {} ~CaffeReluParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeRelu6Parser : public CaffeNodeParser { @@ -36,7 +37,7 @@ class CaffeRelu6Parser : public CaffeNodeParser { CaffeRelu6Parser() : CaffeNodeParser("relu6") {} ~CaffeRelu6Parser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeSigmoidParser : public CaffeNodeParser { @@ -44,7 +45,7 @@ class CaffeSigmoidParser : public CaffeNodeParser { CaffeSigmoidParser() : CaffeNodeParser("sigmoid") {} ~CaffeSigmoidParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeTanhParser : public CaffeNodeParser { @@ -52,7 +53,7 @@ class CaffeTanhParser : public CaffeNodeParser { CaffeTanhParser() : CaffeNodeParser("tanh") {} ~CaffeTanhParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; class CaffeEluParser : public CaffeNodeParser { @@ -60,7 +61,7 @@ class CaffeEluParser : public CaffeNodeParser { CaffeEluParser() : CaffeNodeParser("elu") {} ~CaffeEluParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.cc index f6ee126563..98a0ef2b64 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeArgMaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeArgMaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_keep_dims(true); @@ -39,7 +39,7 @@ ops::PrimitiveC *CaffeArgMaxParser::Parse(const caffe::LayerParameter &proto, co prim->set_axis(argmaxParam.axis()); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeArgMaxParser("ArgMax", new CaffeArgMaxParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.h index 317721d472..44c284ebfe 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_argmax_parser.h @@ -28,7 +28,7 @@ class CaffeArgMaxParser : public CaffeNodeParser { CaffeArgMaxParser() : CaffeNodeParser("argmax") {} ~CaffeArgMaxParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.cc index 5f8a1a6622..fa4d5282f6 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.cc @@ -26,14 +26,15 @@ namespace mindspore { namespace lite { using STATUS = int; -ops::PrimitiveC *CaffeBatchNormParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeBatchNormParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_is_training(false); auto value_ptr = MakeValue(mindspore::Format::NCHW); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, value_ptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, value_ptr); const caffe::BatchNormParameter &batchNormParam = proto.batch_norm_param(); if (proto.bottom_size() != 1) { @@ -53,10 +54,10 @@ ops::PrimitiveC *CaffeBatchNormParser::Parse(const caffe::LayerParameter &proto, } prim->set_epsilon(epsilon); - prim->AddAttr(ops::kUseGlobalStats, MakeValue(true)); + prim_c->AddAttr(ops::kUseGlobalStats, MakeValue(true)); int fmk_type = converter::FmkType::kFmkTypeCaffe; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim_c->AddAttr(ops::kFmkType, MakeValue(fmk_type)); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeBatchNormParser("BatchNorm", new CaffeBatchNormParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.h index f66ed322ff..31f8053720 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_batchnorm_parser.h @@ -28,7 +28,7 @@ class CaffeBatchNormParser : public CaffeNodeParser { CaffeBatchNormParser() : CaffeNodeParser("batchnorm") {} ~CaffeBatchNormParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.cc index 27a8e0f8f1..d624d33a94 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeConcatParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeConcatParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -45,7 +45,7 @@ ops::PrimitiveC *CaffeConcatParser::Parse(const caffe::LayerParameter &proto, co } prim->set_axis(axis); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeConcatParser("Concat", new CaffeConcatParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.h index eef3caee0c..a5e18133dd 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_concat_parser.h @@ -28,7 +28,7 @@ class CaffeConcatParser : public CaffeNodeParser { CaffeConcatParser() : CaffeNodeParser("concat") {} ~CaffeConcatParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.cc index 070db9196f..0930bf1db5 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.cc @@ -21,16 +21,16 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeConvolutionParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeConvolutionParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_pad({0, 0, 0, 0}); prim->set_pad_mode(mindspore::PadMode::PAD); auto value_ptr = MakeValue(mindspore::Format::NCHW); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, value_ptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, value_ptr); prim->set_activation_type(mindspore::NO_ACTIVATION); const caffe::ConvolutionParameter &convParam = proto.convolution_param(); @@ -85,10 +85,10 @@ ops::PrimitiveC *CaffeConvolutionParser::Parse(const caffe::LayerParameter &prot if (group != 1 && group == channel_out) { auto bool_ptr = MakeValue(true); MS_CHECK_TRUE_RET(bool_ptr != nullptr, nullptr); - prim->AddAttr(ops::kIsDepthWise, bool_ptr); + prim_c->AddAttr(ops::kIsDepthWise, bool_ptr); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeConvolutionParser("Convolution", new CaffeConvolutionParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.h index 51a2066d4a..b5ea13e67c 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_convolution_parser.h @@ -29,7 +29,7 @@ class CaffeConvolutionParser : public CaffeNodeParser { CaffeConvolutionParser() : CaffeNodeParser("convolution") {} ~CaffeConvolutionParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.cc index b108fcd670..2c960d436e 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeCropParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeCropParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -50,7 +50,7 @@ ops::PrimitiveC *CaffeCropParser::Parse(const caffe::LayerParameter &proto, cons } } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeCropParser("Crop", new CaffeCropParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.h index e8667940be..01b2629ecd 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_crop_parser.h @@ -28,7 +28,7 @@ class CaffeCropParser : public CaffeNodeParser { CaffeCropParser() : CaffeNodeParser("crop") {} ~CaffeCropParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.cc index e22c5b075a..d5c2c00780 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.cc @@ -22,15 +22,16 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeDeconvolutionParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeDeconvolutionParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim->set_pad({0, 0, 0, 0}); auto value_ptr = MakeValue(mindspore::Format::NCHW); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, value_ptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, value_ptr); prim->set_pad_mode(mindspore::PadMode::PAD); prim->set_output_paddings({0, 0}); @@ -87,14 +88,14 @@ ops::PrimitiveC *CaffeDeconvolutionParser::Parse(const caffe::LayerParameter &pr if (group != 1) { auto bool_ptr = MakeValue(true); MS_CHECK_TRUE_RET(bool_ptr != nullptr, nullptr); - prim->AddAttr(ops::kIsDepthWise, bool_ptr); + prim_c->AddAttr(ops::kIsDepthWise, bool_ptr); } int fmk_type = converter::FmkType::kFmkTypeCaffe; auto fmk_type_ptr = MakeValue(fmk_type); MS_CHECK_TRUE_RET(fmk_type_ptr != nullptr, nullptr); - prim->AddAttr(ops::kFmkType, fmk_type_ptr); - return prim.release(); + prim_c->AddAttr(ops::kFmkType, fmk_type_ptr); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeDeconvolutionParser("Deconvolution", new CaffeDeconvolutionParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.h index 2d9c88a4a5..f5b7acd593 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_deconvolution_parser.h @@ -29,7 +29,7 @@ class CaffeDeconvolutionParser : public CaffeNodeParser { CaffeDeconvolutionParser() : CaffeNodeParser("deconvolution") {} ~CaffeDeconvolutionParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.cc index 69b5938837..039cd12018 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeEltwiseParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeEltwiseParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -69,7 +69,7 @@ ops::PrimitiveC *CaffeEltwiseParser::Parse(const caffe::LayerParameter &proto, c prim->set_mode(mindspore::EltwiseMode::SUM); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeEltwiseParser("Eltwise", new CaffeEltwiseParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.h index 13ededeeb8..e25f4747b2 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_eltwise_parser.h @@ -28,7 +28,7 @@ class CaffeEltwiseParser : public CaffeNodeParser { CaffeEltwiseParser() : CaffeNodeParser("eltwise") {} ~CaffeEltwiseParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.cc index 8f3695f326..1c4a33ac6b 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeExpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeExpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -43,7 +43,7 @@ ops::PrimitiveC *CaffeExpParser::Parse(const caffe::LayerParameter &proto, const prim->set_shift(0); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeExpParser("Exp", new CaffeExpParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.h index c6e649e30e..e9041c27b8 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_exp_parser.h @@ -28,7 +28,7 @@ class CaffeExpParser : public CaffeNodeParser { CaffeExpParser() : CaffeNodeParser("exp") {} ~CaffeExpParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.cc index f89ea2a071..814ba4ed10 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.cc @@ -21,11 +21,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeFlattenParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeFlattenParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_CaffeFlattenParser("Flatten", new CaffeFlattenParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.h index 93b3d4ea27..629cdbcdbd 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_flatten_parser.h @@ -27,7 +27,7 @@ class CaffeFlattenParser : public CaffeNodeParser { CaffeFlattenParser() : CaffeNodeParser("flatten") {} ~CaffeFlattenParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace mindspore::lite diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.cc index a09ba44706..d91b22a825 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.cc @@ -22,11 +22,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeInnerProductParser::Parse(const caffe::LayerParameter &proto, - const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeInnerProductParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::NO_ACTIVATION); const caffe::InnerProductParameter &innerProductParam = proto.inner_product_param(); @@ -37,7 +37,7 @@ ops::PrimitiveC *CaffeInnerProductParser::Parse(const caffe::LayerParameter &pro int64_t num_output = static_cast(innerProductParam.num_output()); auto value_ptr = MakeValue(num_output); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr(ops::kNumOutput, value_ptr); + prim_c->AddAttr(ops::kNumOutput, value_ptr); if (innerProductParam.axis() == 1) { prim->set_axis(1); @@ -50,7 +50,7 @@ ops::PrimitiveC *CaffeInnerProductParser::Parse(const caffe::LayerParameter &pro prim->set_has_bias(true); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeInnerProductParser("InnerProduct", new CaffeInnerProductParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.h index c02193a99a..9ed5e5f64e 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_innerproduct_parser.h @@ -28,7 +28,7 @@ class CaffeInnerProductParser : public CaffeNodeParser { CaffeInnerProductParser() : CaffeNodeParser("innerproduct") {} ~CaffeInnerProductParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.cc index 5bead2a22d..368892b0dd 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.cc @@ -21,10 +21,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeInterpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeInterpParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim->set_method(mindspore::ResizeMethod::LINEAR); prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::ALIGN_CORNERS); @@ -48,9 +50,9 @@ ops::PrimitiveC *CaffeInterpParser::Parse(const caffe::LayerParameter &proto, co } if (interp_param.has_zoom_factor()) { - prim->AddAttr("zoom_factor", MakeValue(interp_param.zoom_factor())); + prim_c->AddAttr("zoom_factor", MakeValue(interp_param.zoom_factor())); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeInterpParser("Interp", new CaffeInterpParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.h index b289b60d96..cfdc69a720 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_interp_parser.h @@ -28,7 +28,7 @@ class CaffeInterpParser : public CaffeNodeParser { CaffeInterpParser() : CaffeNodeParser("Interp") {} ~CaffeInterpParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.cc index 0997463243..5a5ddcfbfa 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.cc @@ -85,6 +85,11 @@ STATUS CheckCaffeModel(const caffe::NetParameter &caffe_model, const caffe::NetP } return RET_OK; } + +FuncGraphPtr ConvertGraph(api::FuncGraphPtr func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} } // namespace bool IsSkipedLayer(const caffe::LayerParameter &layer) { if (layer.type() == "Input" || layer.type() == "Dropout" || layer.type() == "Split") { @@ -134,7 +139,9 @@ api::FuncGraphPtr CaffeModelParser::Parse(const converter::ConverterParameters & ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; } - res_graph_ = std::make_shared(); + auto graph = std::make_shared(); + MS_CHECK_TRUE_MSG(graph != nullptr, nullptr, "create FuncGraph failed"); + res_graph_ = api::MakeShared(graph); MS_CHECK_TRUE_RET(res_graph_ != nullptr, nullptr); status = ConvertGraphInputs(); if (status != RET_OK) { @@ -153,20 +160,18 @@ api::FuncGraphPtr CaffeModelParser::Parse(const converter::ConverterParameters & ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; } - res_graph_->set_attr("graph_name", MakeValue("main_graph")); + graph->set_attr("graph_name", MakeValue("main_graph")); auto value_ptr = MakeValue(static_cast(converter::kFmkTypeCaffe)); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - res_graph_->set_attr("fmk", value_ptr); - auto func_graph = std::dynamic_pointer_cast(res_graph_); - MS_CHECK_TRUE_RET(func_graph != nullptr, nullptr); - if ((status = CommonAnfAdjust(func_graph)) != RET_OK) { + graph->set_attr("fmk", value_ptr); + if ((status = CommonAnfAdjust(graph)) != RET_OK) { MS_LOG(ERROR) << "AdjustForAnf failed."; ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; } auto unify_format = std::make_shared(kFmkTypeCaffe, false); MS_CHECK_TRUE_RET(unify_format != nullptr, nullptr); - if (!unify_format->Run(func_graph)) { + if (!unify_format->Run(graph)) { MS_LOG(ERROR) << "Run insert transpose failed."; return nullptr; } @@ -198,10 +203,10 @@ STATUS CaffeModelParser::ConvertLayers() { // parse primitive MS_LOG(INFO) << "parse op : " << layer.type(); - ops::PrimitiveC *primitive_c = nullptr; + ops::PrimitiveCPtr primitive_c; auto node_parser = registry::NodeParserRegistry::GetNodeParser(kFmkTypeCaffe, layer.type()); if (node_parser != nullptr) { - primitive_c = node_parser->Parse(layer, weight); + primitive_c = node_parser->Parse(layer, weight)->GetPrim(); } else { auto node_parser_builtin = CaffeNodeParserRegistry::GetInstance()->GetNodeParser(layer.type()); if (node_parser_builtin == nullptr) { @@ -237,12 +242,14 @@ STATUS CaffeModelParser::ConvertLayers() { } // build cnode - auto value_node = NewValueNode(std::shared_ptr(primitive_c)); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto value_node = NewValueNode(primitive_c); MSLITE_CHECK_PTR(value_node); std::vector op_inputs = {value_node}; op_inputs.insert(op_inputs.end(), input_nodes.begin(), input_nodes.end()); op_inputs.insert(op_inputs.end(), const_parameters.begin(), const_parameters.end()); - auto new_cnode = res_graph_->NewCNode(op_inputs); + auto new_cnode = graph->NewCNode(op_inputs); MSLITE_CHECK_PTR(new_cnode); new_cnode->set_fullname_with_scope(layer.name()); @@ -307,7 +314,9 @@ STATUS CaffeModelParser::ConvertGraphInputsOfLayer() { MS_LOG(ERROR) << "The input layer should not have inputs"; return RET_ERROR; } - auto parameter = res_graph_->add_parameter(); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto parameter = graph->add_parameter(); MSLITE_CHECK_PTR(parameter); std::vector shape = ConverterInnerContext::GetInstance()->GetGraphInputTensorShape(layer.name()); if (ConverterInnerContext::GetInstance()->GetGraphInputTensorShapeMapSize() > 0 && shape.empty()) { @@ -344,7 +353,9 @@ STATUS CaffeModelParser::ConvertGraphInputsOfShape() { shape_vector.push_back(shape.dim(j)); } } - auto parameter = res_graph_->add_parameter(); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto parameter = graph->add_parameter(); MSLITE_CHECK_PTR(parameter); auto tensor_info = CreateTensorInfo(nullptr, 0, shape_vector, kNumberTypeFloat32); if (tensor_info == nullptr) { @@ -382,7 +393,9 @@ STATUS CaffeModelParser::ConvertGraphInputsOfDim() { } } } - auto parameter = res_graph_->add_parameter(); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto parameter = graph->add_parameter(); auto abstract = CreateTensorAbstract(shape, kNumberTypeFloat32); if (abstract == nullptr) { MS_LOG(ERROR) << "Create tensor abstarct failed"; @@ -420,7 +433,9 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { std::vector make_tuple_inputs; auto make_tuple_prim_ptr = std::make_shared(); MSLITE_CHECK_PTR(make_tuple_prim_ptr); - auto make_tuple_prim = NewValueNode(make_tuple_prim_ptr); + auto make_tuple_prim_c = make_tuple_prim_ptr->GetPrim(); + MSLITE_CHECK_PTR(make_tuple_prim_c); + auto make_tuple_prim = NewValueNode(make_tuple_prim_c); MSLITE_CHECK_PTR(make_tuple_prim); make_tuple_inputs.emplace_back(make_tuple_prim); for (const auto &output_node : caffeInspector.GetGraphOutput()) { @@ -431,25 +446,31 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { auto cnode = nodes_.find(output_node)->second; make_tuple_inputs.emplace_back(cnode); } - auto make_tuple_cnode = res_graph_->NewCNode(make_tuple_inputs); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto make_tuple_cnode = graph->NewCNode(make_tuple_inputs); MSLITE_CHECK_PTR(make_tuple_cnode); make_tuple_cnode->set_fullname_with_scope("return tuple"); std::vector op_inputs; auto return_prim_ptr = std::make_shared(); MSLITE_CHECK_PTR(return_prim_ptr); - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + MSLITE_CHECK_PTR(return_prim_c); + auto value_node = NewValueNode(return_prim_c); MSLITE_CHECK_PTR(value_node); op_inputs.emplace_back(value_node); op_inputs.emplace_back(make_tuple_cnode); - auto cnode = res_graph_->NewCNode(op_inputs); + auto cnode = graph->NewCNode(op_inputs); MSLITE_CHECK_PTR(cnode); cnode->set_fullname_with_scope("Return"); - res_graph_->set_return(cnode); + graph->set_return(cnode); } else { auto returnPrim = std::make_shared(); MSLITE_CHECK_PTR(returnPrim); - auto valueNode = NewValueNode(returnPrim); + auto return_prim_c = returnPrim->GetPrim(); + MSLITE_CHECK_PTR(return_prim_c); + auto valueNode = NewValueNode(return_prim_c); MSLITE_CHECK_PTR(valueNode); std::vector opInputs{valueNode}; if (nodes_.find(caffeInspector.GetGraphOutput().front()) == nodes_.end()) { @@ -462,10 +483,12 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { return RET_NOT_FIND_OP; } opInputs.emplace_back(cnode); - auto returnCnode = res_graph_->NewCNode(opInputs); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto returnCnode = graph->NewCNode(opInputs); MSLITE_CHECK_PTR(returnCnode); returnCnode->set_fullname_with_scope("Return"); - res_graph_->set_return(returnCnode); + graph->set_return(returnCnode); } // save original output tensor names. ConverterInnerContext::GetInstance()->SetGraphOutputTensorNames(caffeInspector.GetGraphOutput()); @@ -473,7 +496,7 @@ STATUS CaffeModelParser::ConvertGraphOutputs() { } STATUS CaffeModelParser::ConvertLayerQuantParams(const caffe::LayerParameter &layer, - const caffe::LayerParameter &weight, ops::PrimitiveC *primitive_c) { + const caffe::LayerParameter &weight, PrimitiveCPtr primitive_c) { MSLITE_CHECK_PTR(primitive_c); auto quant_params_holder = std::make_shared(layer.bottom_size() + weight.blobs_size(), layer.top_size()); @@ -505,7 +528,9 @@ STATUS CaffeModelParser::ConvertBlobs(const caffe::LayerParameter &layer, std::v } // cal Weight num - auto parameter = res_graph_->add_parameter(); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + auto parameter = graph->add_parameter(); MSLITE_CHECK_PTR(parameter); auto type_ptr = TypeIdToType(TypeId::kNumberTypeFloat32); std::vector shape_vector; @@ -591,12 +616,16 @@ STATUS CaffeModelParser::ConvertTop(const caffe::LayerParameter &layer, const CN MS_LOG(ERROR) << "new TupleGetItem failed"; return RET_NULL_PTR; } - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); + auto tuple_get_item_prim_c = tuple_get_item_prim_ptr->GetPrim(); + MSLITE_CHECK_PTR(tuple_get_item_prim_c); + auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_c); MSLITE_CHECK_PTR(tuple_get_item_prim); auto get_item_value = NewValueNode(MakeValue(i)); MSLITE_CHECK_PTR(get_item_value); std::vector inputs{tuple_get_item_prim, cnode, get_item_value}; - CNodePtr get_item_cnode = res_graph_->NewCNode(inputs); + auto graph = ConvertGraph(res_graph_); + MSLITE_CHECK_PTR(graph); + CNodePtr get_item_cnode = graph->NewCNode(inputs); MSLITE_CHECK_PTR(get_item_cnode); get_item_cnode->set_fullname_with_scope(layer.top(i)); nodes_[layer.top(i)] = get_item_cnode; diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.h index 21116d1bd3..8e1f97cbc3 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_model_parser.h @@ -46,7 +46,7 @@ class CaffeModelParser : public converter::ModelParser { STATUS ConvertLayers(); static STATUS ConvertLayerQuantParams(const caffe::LayerParameter &layer, const caffe::LayerParameter &weight, - ops::PrimitiveC *primitive_c); + PrimitiveCPtr primitive_c); STATUS ConvertBlobs(const caffe::LayerParameter &layer, std::vector *const_parameters); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_node_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_node_parser.h index ba736b0d3e..99589ec182 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_node_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_node_parser.h @@ -27,6 +27,7 @@ #include "src/common/log_adapter.h" #include "ops/primitive_c.h" #include "mindspore/core/utils/check_convert_utils.h" +#include "tools/converter/parser/parser_utils.h" namespace mindspore { namespace lite { @@ -36,7 +37,7 @@ class CaffeNodeParser { virtual ~CaffeNodeParser() {} - virtual ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { + virtual PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { return nullptr; } diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.cc index 2fb7ee5efa..034d68909d 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.cc @@ -21,10 +21,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePermuteParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffePermuteParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + std::vector perm; const caffe::PermuteParameter &permuteParam = proto.permute_param(); const int num_order_dims = permuteParam.order_size(); @@ -34,9 +36,9 @@ ops::PrimitiveC *CaffePermuteParser::Parse(const caffe::LayerParameter &proto, c } auto value_ptr = MakeValue(perm); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr("perm", value_ptr); + prim_c->AddAttr("perm", value_ptr); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffePermuteParser("Permute", new CaffePermuteParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.h index 2e230386f3..12dff5b601 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_permute_parser.h @@ -28,7 +28,7 @@ class CaffePermuteParser : public CaffeNodeParser { CaffePermuteParser() : CaffeNodeParser("Permute") {} ~CaffePermuteParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.cc index b200cb827f..9d5a99a65f 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.cc @@ -105,7 +105,7 @@ mindspore::RoundMode CaffePoolingParser::ParseRoundMode(const caffe::PoolingPara return roundMode; } -ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffePoolingParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { const caffe::PoolingParameter &poolingParam = proto.pooling_param(); // parse kernel params @@ -134,7 +134,9 @@ ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, c if (poolingParam.pool() == caffe::PoolingParameter::MAX) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); prim->set_pad_mode(mindspore::PadMode::PAD); prim->set_kernel_size(windows); prim->set_strides(strides); @@ -142,12 +144,14 @@ ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, c prim->set_round_mode(roundMode); prim->set_global(poolingParam.global_pooling()); int fmk_type = converter::FmkType::kFmkTypeCaffe; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim_c->AddAttr(ops::kFmkType, MakeValue(fmk_type)); + return prim->GetPrim(); } else if (poolingParam.pool() == caffe::PoolingParameter::AVE) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); prim->set_pad_mode(mindspore::PadMode::PAD); prim->set_kernel_size(windows); prim->set_strides(strides); @@ -155,8 +159,8 @@ ops::PrimitiveC *CaffePoolingParser::Parse(const caffe::LayerParameter &proto, c prim->set_round_mode(roundMode); prim->set_global(poolingParam.global_pooling()); int fmk_type = converter::kFmkTypeCaffe; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim_c->AddAttr(ops::kFmkType, MakeValue(fmk_type)); + return prim->GetPrim(); } else { MS_LOG(ERROR) << "poolingParam.pool() is not MAX or AVE"; return nullptr; diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.h index c91109e260..b28fbf7059 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_pooling_parser.h @@ -28,7 +28,7 @@ class CaffePoolingParser : public CaffeNodeParser { CaffePoolingParser() : CaffeNodeParser("pooling") {} ~CaffePoolingParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; static STATUS ParsePads(const caffe::PoolingParameter &poolingParam, std::vector *pad); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.cc index df35fc4b99..e4e6a5ca38 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.cc @@ -22,9 +22,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePowerParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffePowerParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); const caffe::PowerParameter &powerParam = proto.power_param(); float power = 1.0; @@ -43,11 +45,11 @@ ops::PrimitiveC *CaffePowerParser::Parse(const caffe::LayerParameter &proto, con } auto value_ptr = MakeValue(power); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr("power", value_ptr); + prim_c->AddAttr("power", value_ptr); prim->set_scale(scale); prim->set_shift(shift); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffePowerParser("Power", new CaffePowerParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.h index 3e320cbb7d..ad57193352 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_power_parser.h @@ -28,7 +28,7 @@ class CaffePowerParser : public CaffeNodeParser { CaffePowerParser() : CaffeNodeParser("power") {} ~CaffePowerParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.cc index 0b79683e45..916fe7d126 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.cc @@ -5,7 +5,7 @@ * 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.0f + * 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, @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffePReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffePReluParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -32,7 +32,7 @@ ops::PrimitiveC *CaffePReluParser::Parse(const caffe::LayerParameter &proto, con prim->set_channel_shared(false); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffePReluParser("PReLU", new CaffePReluParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.h index e9e2669dd0..e18fac1424 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_prelu_parser.h @@ -28,7 +28,7 @@ class CaffePReluParser : public CaffeNodeParser { CaffePReluParser() : CaffeNodeParser("pRelu") {} ~CaffePReluParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.cc index 116ae91008..ef7b65d1dc 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.cc @@ -22,10 +22,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReduceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeReduceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_keep_dims(false); prim->set_reduce_to_end(true); @@ -55,9 +56,9 @@ ops::PrimitiveC *CaffeReduceParser::Parse(const caffe::LayerParameter &proto, co } auto value_ptr = MakeValue(axes); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr("axes", value_ptr); + prim_c->AddAttr("axes", value_ptr); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeReduceParser("Reduction", new CaffeReduceParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.h index ff87638be4..4594c90096 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_reduce_parser.h @@ -28,7 +28,7 @@ class CaffeReduceParser : public CaffeNodeParser { CaffeReduceParser() : CaffeNodeParser("reduce") {} ~CaffeReduceParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.cc index 9d3cbc2856..1d01811eeb 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.cc @@ -21,10 +21,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeReshapeParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeReshapeParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); const caffe::ReshapeParameter &reshapeParam = proto.reshape_param(); if (!reshapeParam.has_shape()) { MS_LOG(ERROR) << "Reshape has no shape info, ret fail"; @@ -37,9 +38,9 @@ ops::PrimitiveC *CaffeReshapeParser::Parse(const caffe::LayerParameter &proto, c } auto value_ptr = MakeValue(shape); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr("shape", value_ptr); + prim_c->AddAttr("shape", value_ptr); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeReshapeParser("Reshape", new CaffeReshapeParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.h index 1456e3d560..acc136709e 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_reshape_parser.h @@ -28,7 +28,7 @@ class CaffeReshapeParser : public CaffeNodeParser { CaffeReshapeParser() : CaffeNodeParser("reshape") {} ~CaffeReshapeParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.cc index d61085cb3b..16bd947547 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.cc @@ -37,7 +37,7 @@ STATUS CaffeScaleParser::GetAxisIndex(const int32_t &axis, uint32_t *axis_index) return RET_OK; } -ops::PrimitiveC *CaffeScaleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeScaleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -58,7 +58,7 @@ ops::PrimitiveC *CaffeScaleParser::Parse(const caffe::LayerParameter &proto, con prim->set_axis(axis_index); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeScaleParser("Scale", new CaffeScaleParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.h index 12cb209215..2ef6d77962 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_scale_parser.h @@ -28,7 +28,7 @@ class CaffeScaleParser : public CaffeNodeParser { CaffeScaleParser() : CaffeNodeParser("scale") {} ~CaffeScaleParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; private: STATUS GetAxisIndex(const int32_t &axis, uint32_t *axis_index); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.cc index fd207bc171..7be841c750 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeSliceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeSliceParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -47,7 +47,7 @@ ops::PrimitiveC *CaffeSliceParser::Parse(const caffe::LayerParameter &proto, con prim->set_axis(slice_param.slice_dim()); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeSliceParser("Slice", new CaffeSliceParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.h index 818a48fa6f..b39299c6db 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_slice_parser.h @@ -28,7 +28,7 @@ class CaffeSliceParser : public CaffeNodeParser { CaffeSliceParser() : CaffeNodeParser("slice") {} ~CaffeSliceParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.cc index ea39b3fe03..0bba1ef928 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeSoftmaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeSoftmaxParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -34,7 +34,7 @@ ops::PrimitiveC *CaffeSoftmaxParser::Parse(const caffe::LayerParameter &proto, c prim->set_axis({1}); } - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeSoftmaxParser("Softmax", new CaffeSoftmaxParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.h index ffe75ec92e..ba52e0152e 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_softmax_parser.h @@ -28,7 +28,7 @@ class CaffeSoftmaxParser : public CaffeNodeParser { CaffeSoftmaxParser() : CaffeNodeParser("softmax") {} ~CaffeSoftmaxParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.cc index 9c73889eb5..3677b75bee 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.cc @@ -22,9 +22,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeTileParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeTileParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); const caffe::TileParameter &tile_param = proto.tile_param(); std::vector dims; @@ -45,9 +47,9 @@ ops::PrimitiveC *CaffeTileParser::Parse(const caffe::LayerParameter &proto, cons } auto value_ptr = MakeValue(multiples); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr("multiples", value_ptr); + prim_c->AddAttr("multiples", value_ptr); - return prim.release(); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeTileParser("Tile", new CaffeTileParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.h index a5f8cfbfaa..ed05bc158d 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_tile_parser.h @@ -28,7 +28,7 @@ class CaffeTileParser : public CaffeNodeParser { CaffeTileParser() : CaffeNodeParser("tile") {} ~CaffeTileParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.cc b/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.cc index 04af70cea3..b0c63214da 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.cc +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.cc @@ -23,10 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { +PrimitiveCPtr CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) { auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_method(mindspore::ResizeMethod::NEAREST); prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::ASYMMETRIC); const caffe::UpsampleParameter &upsample_param = proto.upsample_param(); @@ -37,10 +38,10 @@ ops::PrimitiveC *CaffeUpsampleParser::Parse(const caffe::LayerParameter &proto, return nullptr; } std::vector scales = {1, scale, scale, 1}; - prim->AddAttr("scale", MakeValue(scales)); + prim_c->AddAttr("scale", MakeValue(scales)); } - prim->AddAttr(ops::kOriginalOpName, MakeValue("Upsample")); - return prim.release(); + prim_c->AddAttr(ops::kOriginalOpName, MakeValue("Upsample")); + return prim->GetPrim(); } CaffeNodeRegistrar g_caffeUpsampleParser("Upsample", new CaffeUpsampleParser()); diff --git a/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.h b/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.h index 1575f50fee..9769ec4e50 100644 --- a/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.h +++ b/mindspore/lite/tools/converter/parser/caffe/caffe_upsample_parser.h @@ -28,7 +28,7 @@ class CaffeUpsampleParser : public CaffeNodeParser { CaffeUpsampleParser() : CaffeNodeParser("Upsample") {} ~CaffeUpsampleParser() override = default; - ops::PrimitiveC *Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; + PrimitiveCPtr Parse(const caffe::LayerParameter &proto, const caffe::LayerParameter &weight) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/conv1d_inout_adjust.cc b/mindspore/lite/tools/converter/parser/conv1d_inout_adjust.cc index fda2e578dd..a2293b46e2 100644 --- a/mindspore/lite/tools/converter/parser/conv1d_inout_adjust.cc +++ b/mindspore/lite/tools/converter/parser/conv1d_inout_adjust.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/parser/conv1d_inout_adjust.h" #include #include @@ -25,6 +27,7 @@ #include "ops/primitive_c.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::lite { namespace { @@ -39,8 +42,10 @@ CNodePtr Conv1DInOutAdjust::NewUnsqueezeOpNode(const FuncGraphPtr &func_graph, c MS_CHECK_TRUE_MSG(!axis.empty(), nullptr, "axis is empty"); auto unsqueeze_prim = std::make_shared(); MS_CHECK_TRUE_MSG(unsqueeze_prim != nullptr, nullptr, "create unsqueeze failed."); - unsqueeze_prim->set_attr("axis", MakeValue(axis)); - ValueNodePtr value_node = NewValueNode(unsqueeze_prim); + auto unsqueeze_prim_c = unsqueeze_prim->GetPrim(); + MS_CHECK_TRUE_MSG(unsqueeze_prim_c != nullptr, nullptr, "create unsqueeze_prim_c failed."); + unsqueeze_prim_c->set_attr("axis", MakeValue(axis)); + ValueNodePtr value_node = NewValueNode(unsqueeze_prim_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "new valueNode failed."); std::vector op_inputs = {value_node, input_node}; auto unsqueeze = func_graph->NewCNode(op_inputs); @@ -53,8 +58,10 @@ CNodePtr Conv1DInOutAdjust::NewSqueezeOpNode(const FuncGraphPtr &func_graph, con const std::vector &axis) { auto squeeze_prim = std::make_shared(); MS_CHECK_TRUE_MSG(squeeze_prim != nullptr, nullptr, "create squeeze failed."); - squeeze_prim->set_attr("axis", MakeValue(axis)); - ValueNodePtr value_node = NewValueNode(squeeze_prim); + auto squeeze_prim_c = squeeze_prim->GetPrim(); + MS_CHECK_TRUE_MSG(squeeze_prim_c != nullptr, nullptr, "create squeeze_prim_c failed."); + squeeze_prim_c->set_attr("axis", MakeValue(axis)); + ValueNodePtr value_node = NewValueNode(squeeze_prim_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "new valueNode failed."); std::vector op_inputs = {value_node, input_node}; auto squeeze = func_graph->NewCNode(op_inputs); @@ -110,17 +117,17 @@ bool Conv1DInOutAdjust::Run(const FuncGraphPtr &func_graph) { !opt::CheckPrimitiveType(cnode, prim::kPrimConv2dTransposeFusion)) { continue; } - auto conv2d_prim = GetValueNode(cnode->input(0)); + auto conv2d_prim = ops::GetOperator(cnode->input(0)); MS_CHECK_TRUE_MSG(conv2d_prim != nullptr, false, "conv2d is nullptr."); MS_CHECK_TRUE_MSG(conv2d_prim->GetAttr(ops::kOriginalFormat) != nullptr, false, "The format of conv2d is nullptr."); std::vector axis; switch (Format(GetValue(conv2d_prim->GetAttr(ops::kOriginalFormat)))) { case mindspore::Format::NWC: - (void)conv2d_prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::NHWC)); + (void)conv2d_prim->AddAttr(mindspore::ops::kOriginalFormat, api::MakeValue(mindspore::NHWC)); axis = {1}; break; case mindspore::Format::NCW: - (void)conv2d_prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::NCHW)); + (void)conv2d_prim->AddAttr(mindspore::ops::kOriginalFormat, api::MakeValue(mindspore::NCHW)); axis = {2}; break; default: diff --git a/mindspore/lite/tools/converter/parser/lstm_adjust_pass.cc b/mindspore/lite/tools/converter/parser/lstm_adjust_pass.cc index c0fbc04ba2..a984fb7424 100644 --- a/mindspore/lite/tools/converter/parser/lstm_adjust_pass.cc +++ b/mindspore/lite/tools/converter/parser/lstm_adjust_pass.cc @@ -13,7 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ - +#define USE_DEPRECATED_API #include "tools/converter/parser/lstm_adjust_pass.h" #include "ops/lstm.h" #include "ops/reshape.h" @@ -139,7 +139,7 @@ int ReplaceLstmNode(const FuncGraphManagerPtr &manager, const FuncGraphPtr &func return RET_ERROR; } MS_CHECK_TRUE_MSG(lstm_cnode->input(0) != nullptr, RET_ERROR, "lstm_cnode->input(0) is nullptr."); - auto primitive_c = GetValueNode>(lstm_cnode->input(0)); + auto primitive_c = ops::GetOperator(lstm_cnode->input(0)); if (primitive_c == nullptr) { return RET_ERROR; } @@ -279,15 +279,19 @@ bool LstmAdjustPass::Run(const FuncGraphPtr &func_graph) { auto get_item_name = cnode->fullname_with_scope(); auto transpose_prim = std::make_shared(); MS_CHECK_TRUE_MSG(transpose_prim != nullptr, false, "transpose_prim is nullptr."); + auto transpose_prim_c = transpose_prim->GetPrim(); + MS_CHECK_TRUE_MSG(transpose_prim_c != nullptr, false, "transpose_prim_c is nullptr."); std::vector perm_value = {0, kOutputBatchIndex, 1, kOutputHiddenIndex}; auto transpose_perm = BuildIntVecParameterNode(func_graph, perm_value, "transpose_" + get_item_name + "_perm"); - auto new_transpose_node = func_graph->NewCNode(transpose_prim, {cnode, transpose_perm}); + auto new_transpose_node = func_graph->NewCNode(transpose_prim_c, {cnode, transpose_perm}); MS_CHECK_TRUE_MSG(new_transpose_node != nullptr, false, "New transpose node failed."); auto reshape_prim = std::make_shared(); MS_CHECK_TRUE_MSG(reshape_prim != nullptr, false, "reshape_prim is nullptr."); + auto reshape_prim_c = reshape_prim->GetPrim(); + MS_CHECK_TRUE_MSG(reshape_prim_c != nullptr, false, "reshape_prim_c is nullptr."); auto reshape_perm = BuildIntVecParameterNode(func_graph, {0, 0, -1}, "reshape_" + get_item_name + "_perm"); - auto new_reshape_node = func_graph->NewCNode(reshape_prim, {new_transpose_node, reshape_perm}); + auto new_reshape_node = func_graph->NewCNode(reshape_prim_c, {new_transpose_node, reshape_perm}); MS_CHECK_TRUE_MSG(new_reshape_node != nullptr, false, "New reshape node failed."); (void)manager->Replace(cnode, new_reshape_node); } diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.cc index 263534a1be..fc4d4ad3ab 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.cc @@ -26,17 +26,17 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxReluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxReluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::RELU); prim->set_min_val(0); prim->set_max_val(FLT_MAX); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxLeakyReluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLeakyReluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -48,10 +48,10 @@ ops::PrimitiveC *OnnxLeakyReluParser::Parse(const onnx::GraphProto &onnx_graph, prim->set_activation_type(mindspore::ActivationType::LEAKY_RELU); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxPReluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxPReluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); MS_CHECK_GE(onnx_node.input_size(), kInputSize1, nullptr); @@ -106,10 +106,10 @@ ops::PrimitiveC *OnnxPReluParser::Parse(const onnx::GraphProto &onnx_graph, cons MS_LOG(WARNING) << "The slope pf prelu is null, which may cause errors."; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxEluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxEluParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::ELU); @@ -120,39 +120,39 @@ ops::PrimitiveC *OnnxEluParser::Parse(const onnx::GraphProto &onnx_graph, const } } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxTanhParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxTanhParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::TANH); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxSigmoidParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSigmoidParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::SIGMOID); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxHardSigmoidParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxHardSigmoidParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::HSIGMOID); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxSoftPlusParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSoftPlusParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::SOFTPLUS); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxReluParser("Relu", new OnnxReluParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.h index 98b2a0f1c9..6c0c13eee9 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_activation_parser.h @@ -16,6 +16,7 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_PARSER_ONNX_RELU_PARSER_H #define MINDSPORE_LITE_TOOLS_CONVERTER_PARSER_ONNX_RELU_PARSER_H +#define USE_DEPRECATED_API #include "tools/converter/parser/onnx/onnx_node_parser.h" #include "tools/converter/parser/onnx/onnx_node_parser_registry.h" @@ -27,7 +28,7 @@ class OnnxReluParser : public OnnxNodeParser { OnnxReluParser() : OnnxNodeParser("Relu") {} ~OnnxReluParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxLeakyReluParser : public OnnxNodeParser { @@ -35,7 +36,7 @@ class OnnxLeakyReluParser : public OnnxNodeParser { OnnxLeakyReluParser() : OnnxNodeParser("LeakyRelu") {} ~OnnxLeakyReluParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxPReluParser : public OnnxNodeParser { @@ -43,7 +44,7 @@ class OnnxPReluParser : public OnnxNodeParser { OnnxPReluParser() : OnnxNodeParser("Prelu") {} ~OnnxPReluParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxEluParser : public OnnxNodeParser { @@ -51,7 +52,7 @@ class OnnxEluParser : public OnnxNodeParser { OnnxEluParser() : OnnxNodeParser("Elu") {} ~OnnxEluParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxTanhParser : public OnnxNodeParser { @@ -59,7 +60,7 @@ class OnnxTanhParser : public OnnxNodeParser { OnnxTanhParser() : OnnxNodeParser("Tanh") {} ~OnnxTanhParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxSigmoidParser : public OnnxNodeParser { @@ -67,7 +68,7 @@ class OnnxSigmoidParser : public OnnxNodeParser { OnnxSigmoidParser() : OnnxNodeParser("Sigmoid") {} ~OnnxSigmoidParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxHardSigmoidParser : public OnnxNodeParser { @@ -75,7 +76,7 @@ class OnnxHardSigmoidParser : public OnnxNodeParser { OnnxHardSigmoidParser() : OnnxNodeParser("HardSigmoid") {} ~OnnxHardSigmoidParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxSoftPlusParser : public OnnxNodeParser { @@ -83,7 +84,7 @@ class OnnxSoftPlusParser : public OnnxNodeParser { OnnxSoftPlusParser() : OnnxNodeParser("Softplus") {} ~OnnxSoftPlusParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.cc index c1e9f13f62..9412aa6deb 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.cc @@ -21,10 +21,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxAdderParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAdderParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxAdderParser("adder_f", new OnnxAdderParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.h index 31f6b131c7..779c1a3d11 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_adder_parser.h @@ -27,7 +27,7 @@ class OnnxAdderParser : public OnnxNodeParser { OnnxAdderParser() : OnnxNodeParser("Adder") {} ~OnnxAdderParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.cc index 85ebffd029..f12734f259 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxArgMaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxArgMaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -33,7 +33,7 @@ ops::PrimitiveC *OnnxArgMaxParser::Parse(const onnx::GraphProto &onnx_graph, con } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxArgMaxParser("ArgMax", new OnnxArgMaxParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.h index 4dea29b724..e23bb67beb 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_argmax_parser.h @@ -27,7 +27,7 @@ class OnnxArgMaxParser : public OnnxNodeParser { OnnxArgMaxParser() : OnnxNodeParser("ArgMax") {} ~OnnxArgMaxParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.cc index 76b0fbb1d7..e3c53095c3 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.cc @@ -48,61 +48,61 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxAddParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAddParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxSubParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSubParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxDivParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxDivParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxMulParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxMulParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxEqualParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxEqualParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxLessParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLessParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxGreaterParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxGreaterParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxFloorParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxFloorParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxAbsParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAbsParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxExpParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxExpParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -110,110 +110,110 @@ ops::PrimitiveC *OnnxExpParser::Parse(const onnx::GraphProto &onnx_graph, const prim->set_scale(1.0); prim->set_shift(0.0); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxCosParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxCosParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxCeilParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxCeilParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxLogParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLogParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxAtanParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAtanParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxAsinParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAsinParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxAndParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAndParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxOrParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxOrParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxNotParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxNotParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxNegParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxNegParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxRoundParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxRoundParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxSinParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSinParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxTanParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxTanParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxSqrtParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSqrtParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxPowParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxPowParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_scale(1.0); prim->set_shift(0.0); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxMinParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxMinParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxMaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxMaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxEltwiseParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxEltwiseParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -224,13 +224,13 @@ ops::PrimitiveC *OnnxEltwiseParser::Parse(const onnx::GraphProto &onnx_graph, co return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxReciprocalParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxReciprocalParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxAddParser("Add", new OnnxAddParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.h index 557c91bcb4..f059fbfed1 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_arithmetic_operation_parser.h @@ -27,7 +27,7 @@ class OnnxAddParser : public OnnxNodeParser { OnnxAddParser() : OnnxNodeParser("Add") {} ~OnnxAddParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxSubParser : public OnnxNodeParser { @@ -35,7 +35,7 @@ class OnnxSubParser : public OnnxNodeParser { OnnxSubParser() : OnnxNodeParser("Sub") {} ~OnnxSubParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxMulParser : public OnnxNodeParser { @@ -43,7 +43,7 @@ class OnnxMulParser : public OnnxNodeParser { OnnxMulParser() : OnnxNodeParser("Mul") {} ~OnnxMulParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxDivParser : public OnnxNodeParser { @@ -51,7 +51,7 @@ class OnnxDivParser : public OnnxNodeParser { OnnxDivParser() : OnnxNodeParser("Div") {} ~OnnxDivParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxPowParser : public OnnxNodeParser { @@ -59,7 +59,7 @@ class OnnxPowParser : public OnnxNodeParser { OnnxPowParser() : OnnxNodeParser("Power") {} ~OnnxPowParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxEqualParser : public OnnxNodeParser { @@ -67,7 +67,7 @@ class OnnxEqualParser : public OnnxNodeParser { OnnxEqualParser() : OnnxNodeParser("Equal") {} ~OnnxEqualParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxLessParser : public OnnxNodeParser { @@ -75,7 +75,7 @@ class OnnxLessParser : public OnnxNodeParser { OnnxLessParser() : OnnxNodeParser("Less") {} ~OnnxLessParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxGreaterParser : public OnnxNodeParser { @@ -83,7 +83,7 @@ class OnnxGreaterParser : public OnnxNodeParser { OnnxGreaterParser() : OnnxNodeParser("Greater") {} ~OnnxGreaterParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxMinParser : public OnnxNodeParser { @@ -91,7 +91,7 @@ class OnnxMinParser : public OnnxNodeParser { OnnxMinParser() : OnnxNodeParser("Min") {} ~OnnxMinParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxMaxParser : public OnnxNodeParser { @@ -99,7 +99,7 @@ class OnnxMaxParser : public OnnxNodeParser { OnnxMaxParser() : OnnxNodeParser("Max") {} ~OnnxMaxParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxEltwiseParser : public OnnxNodeParser { @@ -107,7 +107,7 @@ class OnnxEltwiseParser : public OnnxNodeParser { OnnxEltwiseParser() : OnnxNodeParser("Eltwise") {} ~OnnxEltwiseParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxFloorParser : public OnnxNodeParser { @@ -115,7 +115,7 @@ class OnnxFloorParser : public OnnxNodeParser { OnnxFloorParser() : OnnxNodeParser("Floor") {} ~OnnxFloorParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxAbsParser : public OnnxNodeParser { @@ -123,7 +123,7 @@ class OnnxAbsParser : public OnnxNodeParser { OnnxAbsParser() : OnnxNodeParser("Abs") {} ~OnnxAbsParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxNegParser : public OnnxNodeParser { @@ -131,7 +131,7 @@ class OnnxNegParser : public OnnxNodeParser { OnnxNegParser() : OnnxNodeParser("Neg") {} ~OnnxNegParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxExpParser : public OnnxNodeParser { @@ -139,7 +139,7 @@ class OnnxExpParser : public OnnxNodeParser { OnnxExpParser() : OnnxNodeParser("Exp") {} ~OnnxExpParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxCosParser : public OnnxNodeParser { @@ -147,7 +147,7 @@ class OnnxCosParser : public OnnxNodeParser { OnnxCosParser() : OnnxNodeParser("Cos") {} ~OnnxCosParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxSinParser : public OnnxNodeParser { @@ -155,7 +155,7 @@ class OnnxSinParser : public OnnxNodeParser { OnnxSinParser() : OnnxNodeParser("Sin") {} ~OnnxSinParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxSqrtParser : public OnnxNodeParser { @@ -163,7 +163,7 @@ class OnnxSqrtParser : public OnnxNodeParser { OnnxSqrtParser() : OnnxNodeParser("Sqrt") {} ~OnnxSqrtParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxCeilParser : public OnnxNodeParser { @@ -171,7 +171,7 @@ class OnnxCeilParser : public OnnxNodeParser { OnnxCeilParser() : OnnxNodeParser("Ceil") {} ~OnnxCeilParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxLogParser : public OnnxNodeParser { @@ -179,7 +179,7 @@ class OnnxLogParser : public OnnxNodeParser { OnnxLogParser() : OnnxNodeParser("Log") {} ~OnnxLogParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxTanParser : public OnnxNodeParser { @@ -187,7 +187,7 @@ class OnnxTanParser : public OnnxNodeParser { OnnxTanParser() : OnnxNodeParser("Tan") {} ~OnnxTanParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxAtanParser : public OnnxNodeParser { @@ -195,7 +195,7 @@ class OnnxAtanParser : public OnnxNodeParser { OnnxAtanParser() : OnnxNodeParser("Atan") {} ~OnnxAtanParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxAsinParser : public OnnxNodeParser { @@ -203,7 +203,7 @@ class OnnxAsinParser : public OnnxNodeParser { OnnxAsinParser() : OnnxNodeParser("Asin") {} ~OnnxAsinParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxAndParser : public OnnxNodeParser { @@ -211,7 +211,7 @@ class OnnxAndParser : public OnnxNodeParser { OnnxAndParser() : OnnxNodeParser("And") {} ~OnnxAndParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxOrParser : public OnnxNodeParser { @@ -219,7 +219,7 @@ class OnnxOrParser : public OnnxNodeParser { OnnxOrParser() : OnnxNodeParser("Or") {} ~OnnxOrParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxNotParser : public OnnxNodeParser { @@ -227,7 +227,7 @@ class OnnxNotParser : public OnnxNodeParser { OnnxNotParser() : OnnxNodeParser("Not") {} ~OnnxNotParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxRoundParser : public OnnxNodeParser { @@ -235,7 +235,7 @@ class OnnxRoundParser : public OnnxNodeParser { OnnxRoundParser() : OnnxNodeParser("Round") {} ~OnnxRoundParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxReciprocalParser : public OnnxNodeParser { @@ -243,7 +243,7 @@ class OnnxReciprocalParser : public OnnxNodeParser { OnnxReciprocalParser() : OnnxNodeParser("Reciprocal") {} ~OnnxReciprocalParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.cc index 1c7493fd66..1d52cc8682 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxBatchNormParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxBatchNormParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -32,7 +32,7 @@ ops::PrimitiveC *OnnxBatchNormParser::Parse(const onnx::GraphProto &onnx_graph, } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxBatchNormParser("BatchNormalization", new OnnxBatchNormParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.h index fff6fcd4a2..44972835ea 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_batchnorm_parser.h @@ -27,7 +27,7 @@ class OnnxBatchNormParser : public OnnxNodeParser { OnnxBatchNormParser() : OnnxNodeParser("BatchNormalization") {} ~OnnxBatchNormParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.cc index 399948bd50..87a9ae0d79 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.cc @@ -21,10 +21,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxBiasAddParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxBiasAddParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxBiasAddParser("BiasAdd", new OnnxBiasAddParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.h index 265ff970fe..ca634f96ea 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_biasadd_parser.h @@ -27,7 +27,7 @@ class OnnxBiasAddParser : public OnnxNodeParser { OnnxBiasAddParser() : OnnxNodeParser("BiasAdd") {} ~OnnxBiasAddParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.cc index 57978b0c15..1e7c58631e 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxCastParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxCastParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -35,11 +35,13 @@ ops::PrimitiveC *OnnxCastParser::Parse(const onnx::GraphProto &onnx_graph, const if (dst_type == kNumberTypeFloat64) { dst_type = kNumberTypeFloat32; } - prim->AddAttr("to", MakeValue(static_cast(dst_type))); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr("to", MakeValue(static_cast(dst_type))); } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxCastParser("Cast", new OnnxCastParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.h index 3bf67beb25..b07d64ae31 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_cast_parser.h @@ -27,7 +27,7 @@ class OnnxCastParser : public OnnxNodeParser { OnnxCastParser() : OnnxNodeParser("Cast") {} ~OnnxCastParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.cc index 108278ed5a..d6a7a315e2 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxClipParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxClipParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_min(-FLT_MAX); @@ -35,7 +35,7 @@ ops::PrimitiveC *OnnxClipParser::Parse(const onnx::GraphProto &onnx_graph, const } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxClipParser("Clip", new OnnxClipParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.h index 44c58fe04c..f7b98a8d73 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_clip_parser.h @@ -27,7 +27,7 @@ class OnnxClipParser : public OnnxNodeParser { OnnxClipParser() : OnnxNodeParser("Clip") {} ~OnnxClipParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.cc index a48eea71e1..34d73dc5e2 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxConcatParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxConcatParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -31,7 +31,7 @@ ops::PrimitiveC *OnnxConcatParser::Parse(const onnx::GraphProto &onnx_graph, con } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxConcatParser("Concat", new OnnxConcatParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.h index fc12edd90f..06922b06b4 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_concat_parser.h @@ -27,7 +27,7 @@ class OnnxConcatParser : public OnnxNodeParser { OnnxConcatParser() : OnnxNodeParser("Concat") {} ~OnnxConcatParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.cc index 5db5d6ed58..a5243571b8 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.cc @@ -23,8 +23,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxConstantOfShapeParser::Parse(const onnx::GraphProto &onnx_graph, - const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxConstantOfShapeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); int data_type = 0; @@ -61,7 +60,7 @@ ops::PrimitiveC *OnnxConstantOfShapeParser::Parse(const onnx::GraphProto &onnx_g prim->set_value(values); prim->set_data_type((int64_t)data_type); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxConstantOfShapeParser("ConstantOfShape", new OnnxConstantOfShapeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.h index 2cabe1e0be..e70e666761 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_of_shape_parser.h @@ -27,7 +27,7 @@ class OnnxConstantOfShapeParser : public OnnxNodeParser { OnnxConstantOfShapeParser() : OnnxNodeParser("ConstantOfShape") {} ~OnnxConstantOfShapeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.cc index 9366e30605..2bf3d9dc4a 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.cc @@ -24,7 +24,7 @@ namespace mindspore { namespace lite { -STATUS OnnxConstantParser::AddDataInfoAttr(const onnx::TensorProto &onnx_const_tensor, ops::PrimitiveC *prim) { +STATUS OnnxConstantParser::AddDataInfoAttr(const onnx::TensorProto &onnx_const_tensor, PrimitiveCPtr prim) { MS_ASSERT(prim != nullptr); auto data_type = OnnxNodeParser::GetDataTypeFromOnnx(static_cast(onnx_const_tensor.data_type())); @@ -47,8 +47,8 @@ STATUS OnnxConstantParser::AddDataInfoAttr(const onnx::TensorProto &onnx_const_t return RET_OK; } -ops::PrimitiveC *OnnxConstantParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { - auto prim = std::make_unique(); +PrimitiveCPtr OnnxConstantParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &attr : onnx_node.attribute()) { if (attr.name() == "sparse_value") { @@ -57,7 +57,7 @@ ops::PrimitiveC *OnnxConstantParser::Parse(const onnx::GraphProto &onnx_graph, c } if (attr.name() == "value") { const auto &const_tensor = attr.t(); - if (AddDataInfoAttr(const_tensor, prim.get()) != RET_OK) { + if (AddDataInfoAttr(const_tensor, prim) != RET_OK) { MS_LOG(ERROR) << "add basic attr failed."; return nullptr; } @@ -66,7 +66,7 @@ ops::PrimitiveC *OnnxConstantParser::Parse(const onnx::GraphProto &onnx_graph, c return nullptr; } } - return prim.release(); + return prim; } OnnxNodeRegistrar g_onnxConstantParser("Constant", new OnnxConstantParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.h index 12e10ee564..119564aef3 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_constant_parser.h @@ -27,10 +27,10 @@ class OnnxConstantParser : public OnnxNodeParser { OnnxConstantParser() : OnnxNodeParser("Constant") {} ~OnnxConstantParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; private: - STATUS AddDataInfoAttr(const onnx::TensorProto &onnx_const_tensor, ops::PrimitiveC *prim); + STATUS AddDataInfoAttr(const onnx::TensorProto &onnx_const_tensor, PrimitiveCPtr prim); }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.cc index 39e1a968e2..5e1ac6ab2c 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.cc @@ -21,6 +21,7 @@ #include #include "ops/fusion/conv2d_fusion.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::lite { STATUS GetConvChannel(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node, int64_t group, @@ -116,9 +117,11 @@ STATUS OnnxConvParser::ParseOnnxAttr(const onnx::NodeProto &onnx_node, int64_t * return RET_OK; } -ops::PrimitiveC *OnnxConvParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxConvParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_pad({0, 0, 0, 0}); mindspore::Format format = mindspore::Format::NCHW; mindspore::PadMode pad_mode = mindspore::PadMode::PAD; @@ -131,7 +134,7 @@ ops::PrimitiveC *OnnxConvParser::Parse(const onnx::GraphProto &onnx_graph, const MS_LOG(ERROR) << "Parse onnx attribute failed."; return nullptr; } - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); prim->set_pad_mode(pad_mode); prim->set_group(group); @@ -140,7 +143,7 @@ ops::PrimitiveC *OnnxConvParser::Parse(const onnx::GraphProto &onnx_graph, const return nullptr; } if (conv1d) { - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(NCW)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(NCW)); } prim->set_dilation({1, 1}); if (!dilation.empty()) { @@ -173,10 +176,10 @@ ops::PrimitiveC *OnnxConvParser::Parse(const onnx::GraphProto &onnx_graph, const } if (group == channel_in && channel_in == channel_out) { - prim->AddAttr(ops::kIsDepthWise, MakeValue(true)); + prim_c->AddAttr(ops::kIsDepthWise, MakeValue(true)); } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxConvParser("Conv", new OnnxConvParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.h index ee8c6e5032..2cb8dbf1a2 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_parser.h @@ -30,7 +30,7 @@ class OnnxConvParser : public OnnxConvBaseParser { OnnxConvParser() : OnnxConvBaseParser("Conv") {} ~OnnxConvParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; private: STATUS ParseOnnxAttr(const onnx::NodeProto &onnx_node, int64_t *group, mindspore::Format *format, diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.cc index d4d04b4b9d..37237b516f 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.cc @@ -45,10 +45,12 @@ STATUS OnnxDeConvParser::ParseOnnxAttr(const onnx::NodeProto &onnx_node, int64_t return RET_OK; } -ops::PrimitiveC *OnnxDeConvParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxDeConvParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { MS_CHECK_GE(onnx_node.input_size(), kInputSize1, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_pad({0, 0, 0, 0}); mindspore::PadMode pad_mode = mindspore::PadMode::PAD; std::vector kernel, dilate, stride, pads, output_paddings; @@ -58,7 +60,7 @@ ops::PrimitiveC *OnnxDeConvParser::Parse(const onnx::GraphProto &onnx_graph, con MS_LOG(ERROR) << "Parse onnx attribute failed."; return nullptr; } - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); prim->set_group(group); prim->set_pad_mode(pad_mode); @@ -67,7 +69,7 @@ ops::PrimitiveC *OnnxDeConvParser::Parse(const onnx::GraphProto &onnx_graph, con return nullptr; } if (conv1d) { - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(NCW)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(NCW)); } if (!dilate.empty()) { prim->set_dilation(dilate); @@ -112,11 +114,11 @@ ops::PrimitiveC *OnnxDeConvParser::Parse(const onnx::GraphProto &onnx_graph, con prim->set_out_channel(weight_shape[1] * group); if (group != 1 && weight_shape[1] == 1) { - prim->AddAttr(ops::kIsDepthWise, MakeValue(true)); + prim_c->AddAttr(ops::kIsDepthWise, MakeValue(true)); } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxDeConvParser("ConvTranspose", new OnnxDeConvParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.h index 3cf43e4928..89e2d80e78 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_conv_transpose_parser.h @@ -29,7 +29,7 @@ class OnnxDeConvParser : public OnnxConvBaseParser { OnnxDeConvParser() : OnnxConvBaseParser("DeConv") {} ~OnnxDeConvParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; private: STATUS ParseOnnxAttr(const onnx::NodeProto &onnx_node, int64_t *group, mindspore::PadMode *pad_mode, diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.cc index 1f7db05da1..83efe6d518 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxDepthToSpaceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxDepthToSpaceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -31,7 +31,7 @@ ops::PrimitiveC *OnnxDepthToSpaceParser::Parse(const onnx::GraphProto &onnx_grap } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxDepthToSpaceParser("DepthToSpace", new OnnxDepthToSpaceParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.h index 3b32e96d40..2c880ad4f8 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_depth_to_space_parser.h @@ -27,7 +27,7 @@ class OnnxDepthToSpaceParser : public OnnxNodeParser { OnnxDepthToSpaceParser() : OnnxNodeParser("DepthToSpace") {} ~OnnxDepthToSpaceParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.cc index 1b7a8986bd..0f952b55fe 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxDropoutParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxDropoutParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -31,7 +31,7 @@ ops::PrimitiveC *OnnxDropoutParser::Parse(const onnx::GraphProto &onnx_graph, co } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxDropoutParser("Dropout", new OnnxDropoutParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.h index be6d33da5d..05d8d4be56 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_dropout_parser.h @@ -27,7 +27,7 @@ class OnnxDropoutParser : public OnnxNodeParser { OnnxDropoutParser() : OnnxNodeParser("Dropout") {} ~OnnxDropoutParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.cc index 646b72e7ba..1d751f73fc 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.cc @@ -20,10 +20,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxErfParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxErfParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnx_erf_parser("Erf", new OnnxErfParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.h index bfdeb4d82e..7361a13ace 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_erf_parser.h @@ -26,7 +26,7 @@ class OnnxErfParser : public OnnxNodeParser { OnnxErfParser() : OnnxNodeParser("Erf") {} ~OnnxErfParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.cc index fb1ed7ca9f..7b442bf61d 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxExpandParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxExpandParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); MS_CHECK_GE(onnx_node.input_size(), kInputSize1, nullptr); @@ -59,7 +59,7 @@ ops::PrimitiveC *OnnxExpandParser::Parse(const onnx::GraphProto &onnx_graph, con prim->set_shape(dst_shape); } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxExpandSpaceParser("Expand", new OnnxExpandParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.h index bb43d24b48..b4ce78f38b 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_expand_parser.h @@ -27,7 +27,7 @@ class OnnxExpandParser : public OnnxNodeParser { OnnxExpandParser() : OnnxNodeParser("Expand") {} ~OnnxExpandParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.cc index 494dff8ff0..a0b047d91f 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.cc @@ -21,10 +21,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxFlattenParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxFlattenParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxFlattenParser("Flatten", new OnnxFlattenParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.h index 8211f751f1..5c8843f672 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_flatten_parser.h @@ -27,7 +27,7 @@ class OnnxFlattenParser : public OnnxNodeParser { OnnxFlattenParser() : OnnxNodeParser("Fatten") {} ~OnnxFlattenParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.cc index e9babab561..e7a18547ec 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.cc @@ -21,9 +21,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxGatherParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxGatherParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); int32_t axis = 0; for (const auto &onnx_node_attr : onnx_node.attribute()) { const auto &attribute_name = onnx_node_attr.name(); @@ -31,9 +33,9 @@ ops::PrimitiveC *OnnxGatherParser::Parse(const onnx::GraphProto &onnx_graph, con axis = static_cast(onnx_node_attr.i()); } } - prim->AddAttr("axis", MakeValue(axis)); + prim_c->AddAttr("axis", MakeValue(axis)); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxGatherParser("Gather", new OnnxGatherParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.h index f213c3643c..2e974e7a0b 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_gather_parser.h @@ -27,7 +27,7 @@ class OnnxGatherParser : public OnnxNodeParser { OnnxGatherParser() : OnnxNodeParser("Gather") {} ~OnnxGatherParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.cc index 0daa1e5ace..6d8b2db608 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.cc @@ -26,7 +26,7 @@ namespace mindspore { namespace lite { namespace { -STATUS ParseInt8GivenIntTensorFill(const onnx::NodeProto &onnx_node, ops::PrimitiveC *prim, +STATUS ParseInt8GivenIntTensorFill(const onnx::NodeProto &onnx_node, PrimitiveCPtr prim, const std::vector &shape) { MS_ASSERT(prim != nullptr); int data_count = 1; @@ -55,8 +55,7 @@ STATUS ParseInt8GivenIntTensorFill(const onnx::NodeProto &onnx_node, ops::Primit return RET_OK; } -STATUS ParseInt8GivenTensorFill(const onnx::NodeProto &onnx_node, ops::PrimitiveC *prim, - const std::vector &shape) { +STATUS ParseInt8GivenTensorFill(const onnx::NodeProto &onnx_node, PrimitiveCPtr prim, const std::vector &shape) { MS_ASSERT(prim != nullptr); int data_count = 1; for (size_t i = 0; i < shape.size(); i++) { @@ -80,9 +79,8 @@ STATUS ParseInt8GivenTensorFill(const onnx::NodeProto &onnx_node, ops::Primitive } } // namespace -ops::PrimitiveC *OnnxGivenTensorFillParser::Parse(const onnx::GraphProto &onnx_graph, - const onnx::NodeProto &onnx_node) { - auto prim = std::make_unique(); +PrimitiveCPtr OnnxGivenTensorFillParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); std::vector shape_vector; auto iter = std::find_if(onnx_node.attribute().begin(), onnx_node.attribute().end(), @@ -94,18 +92,18 @@ ops::PrimitiveC *OnnxGivenTensorFillParser::Parse(const onnx::GraphProto &onnx_g std::transform(shape_vector.begin(), shape_vector.end(), std::back_inserter(shape), [](const int64_t &val) { return static_cast(val); }); if (onnx_node.op_type() == "Int8GivenIntTensorFill") { - if (ParseInt8GivenIntTensorFill(onnx_node, prim.get(), shape) != RET_OK) { + if (ParseInt8GivenIntTensorFill(onnx_node, prim, shape) != RET_OK) { MS_LOG(ERROR) << "given tensor fill parse failed."; return nullptr; } } else if (onnx_node.op_type() == "Int8GivenTensorFill") { - if (ParseInt8GivenTensorFill(onnx_node, prim.get(), shape) != RET_OK) { + if (ParseInt8GivenTensorFill(onnx_node, prim, shape) != RET_OK) { MS_LOG(ERROR) << "given tensor fill parse failed."; return nullptr; } } - return prim.release(); + return prim; } OnnxNodeRegistrar g_onnxInt8GivenIntTensorFillParser("Int8GivenIntTensorFill", new OnnxGivenTensorFillParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.h index 7e8d55b6e7..73f78e25ea 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_given_tensor_fill_parser.h @@ -28,7 +28,7 @@ class OnnxGivenTensorFillParser : public OnnxNodeParser { OnnxGivenTensorFillParser() : OnnxNodeParser("GivenTensorFill") {} ~OnnxGivenTensorFillParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.cc index 300c1573e6..7aabb3b75a 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.cc @@ -22,10 +22,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxIdentityParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxIdentityParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxIdentityParser("Identity", new OnnxIdentityParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.h index 4dea7165c1..0bbcd9e171 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_identity_parser.h @@ -27,7 +27,7 @@ class OnnxIdentityParser : public OnnxNodeParser { OnnxIdentityParser() : OnnxNodeParser("Identity") {} ~OnnxIdentityParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.cc index c637f3a070..60066820e6 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.cc @@ -22,10 +22,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxIfParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { - auto prim = std::make_unique(); +PrimitiveCPtr OnnxIfParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim; } OnnxNodeRegistrar g_onnxIfParser("If", new OnnxIfParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.h index 2fe114bfbf..24d45d12f0 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_if_parser.h @@ -27,7 +27,7 @@ class OnnxIfParser : public OnnxNodeParser { OnnxIfParser() : OnnxNodeParser("If") {} ~OnnxIfParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_inputs_adjust.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_inputs_adjust.cc index 7aff49f937..1162c19dc6 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_inputs_adjust.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_inputs_adjust.cc @@ -25,6 +25,7 @@ #include "nnacl/op_base.h" #include "tools/common/tensor_util.h" #include "tools/optimizer/common/gllo_utils.h" +#include "tools/common/node_util.h" namespace mindspore::lite { namespace { @@ -324,7 +325,7 @@ STATUS AdjustResize(bool *need_update_manager, const CNodePtr &cnode) { MS_CHECK_TRUE_RET(!cnode->inputs().empty(), lite::RET_ERROR); auto node = cnode->input(0); MS_ASSERT(node != nullptr); - auto resize_prim = GetValueNode>(node); + auto resize_prim = GetValueNode>(node); if (resize_prim == nullptr) { MS_LOG(ERROR) << "cnode is invalid."; return lite::RET_ERROR; @@ -386,7 +387,9 @@ STATUS AdjustRandomNormal(const FuncGraphPtr &func_graph, const CNodePtr &cnode) if (cnode->size() != 1) { return RET_OK; } - auto prim = GetValueNode>(cnode->input(0)); + auto random_normal_node = ops::GetOperator(cnode->input(0)); + MS_CHECK_TRUE_RET(random_normal_node != nullptr, RET_ERROR); + auto prim = random_normal_node->GetPrim(); MS_CHECK_TRUE_RET(prim != nullptr, RET_ERROR); MS_CHECK_TRUE_RET(prim->GetAttr(ops::kDataType) != nullptr, RET_ERROR); TypeId data_type = static_cast(GetValue(prim->GetAttr(ops::kDataType))); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.cc index 074467e4ab..8211f1bd11 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxInstanceNormParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxInstanceNormParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); if (!onnx_node.attribute().empty()) { @@ -31,7 +31,7 @@ ops::PrimitiveC *OnnxInstanceNormParser::Parse(const onnx::GraphProto &onnx_grap } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxInstanceNormParser("InstanceNormalization", new OnnxInstanceNormParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.h index 155409bd09..bbd2a40d72 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_instance_norm_parser.h @@ -27,7 +27,7 @@ class OnnxInstanceNormParser : public OnnxNodeParser { OnnxInstanceNormParser() : OnnxNodeParser("InstanceNorm") {} ~OnnxInstanceNormParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.cc index 55eb641739..90df28da69 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxLogSoftmaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLogSoftmaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); int64_t axis = -1; @@ -33,7 +33,7 @@ ops::PrimitiveC *OnnxLogSoftmaxParser::Parse(const onnx::GraphProto &onnx_graph, } prim->set_axis(axis); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxLogSoftmaxParser("LogSoftmax", new OnnxLogSoftmaxParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.h index 18569cb00b..d2cf28277c 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_log_softmax_parser.h @@ -27,7 +27,7 @@ class OnnxLogSoftmaxParser : public OnnxNodeParser { OnnxLogSoftmaxParser() : OnnxNodeParser("LogSoftmax") {} ~OnnxLogSoftmaxParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.cc index 2bd4dd3802..dec3f2dd3a 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.cc @@ -22,10 +22,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxLoopParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { - auto prim = std::make_unique(); +PrimitiveCPtr OnnxLoopParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim; } OnnxNodeRegistrar g_onnxLoopParser("Loop", new OnnxLoopParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.h index 700be8baf2..d54fec4fa2 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_loop_parser.h @@ -27,7 +27,7 @@ class OnnxLoopParser : public OnnxNodeParser { OnnxLoopParser() : OnnxNodeParser("Loop") {} ~OnnxLoopParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.cc index fd46827570..4b46eaae16 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxLpNormParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLpNormParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -33,7 +33,7 @@ ops::PrimitiveC *OnnxLpNormParser::Parse(const onnx::GraphProto &onnx_graph, con } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxLpNormParser("LpNormalization", new OnnxLpNormParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.h index 1beef0a78b..5f0c35e016 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_lp_norm_parser.h @@ -27,7 +27,7 @@ class OnnxLpNormParser : public OnnxNodeParser { OnnxLpNormParser() : OnnxNodeParser("LpNorm") {} ~OnnxLpNormParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.cc index 2359dd89c5..fbce0e9f38 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxLrnParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLrnParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); int64_t size = 0; @@ -47,7 +47,7 @@ ops::PrimitiveC *OnnxLrnParser::Parse(const onnx::GraphProto &onnx_graph, const alpha /= size; prim->set_alpha(alpha); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxLrnxParser("Lrn", new OnnxLrnParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.h index 3fae8c0977..38431e0829 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_lrn_parser.h @@ -27,7 +27,7 @@ class OnnxLrnParser : public OnnxNodeParser { OnnxLrnParser() : OnnxNodeParser("Lrn") {} ~OnnxLrnParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.cc index f7050955fe..a60c83ee75 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxLstmParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxLstmParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -43,7 +43,7 @@ ops::PrimitiveC *OnnxLstmParser::Parse(const onnx::GraphProto &onnx_graph, const } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxLstmParser("LSTM", new OnnxLstmParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.h index 5262188065..7c48265e52 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_lstm_parser.h @@ -27,7 +27,7 @@ class OnnxLstmParser : public OnnxNodeParser { OnnxLstmParser() : OnnxNodeParser("LSTM") {} ~OnnxLstmParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.cc index c18d09f510..ad31a212eb 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxMatmulParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxMatmulParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); float alpha = 1.0f; @@ -43,7 +43,7 @@ ops::PrimitiveC *OnnxMatmulParser::Parse(const onnx::GraphProto &onnx_graph, con return nullptr; } prim->set_activation_type(mindspore::NO_ACTIVATION); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxMatmulParser("MatMul", new OnnxMatmulParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.h index 9d9e7ac6fa..0db0050ab4 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_matmul_parser.h @@ -27,7 +27,7 @@ class OnnxMatmulParser : public OnnxNodeParser { OnnxMatmulParser() : OnnxNodeParser("MatMul") {} ~OnnxMatmulParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.cc index 894c039f39..9c060c86b2 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.cc @@ -13,7 +13,6 @@ * See the License for the specific language governing permissions and * limitations under the License. */ - #include "tools/converter/parser/onnx/onnx_model_parser.h" #include #include @@ -159,7 +158,7 @@ CNodePtr GetCNodeFromControlFlowNodesMap( auto iter1 = control_nodes_map.find(loop_node_name); if (iter1 == control_nodes_map.end()) { return nullptr; - } + } // namespace auto iter2 = iter1->second->find(loop_node_name); if (iter2 == iter1->second->end()) { return nullptr; @@ -174,7 +173,9 @@ STATUS BuildReturnNode(const FuncGraphPtr &anf_graph, const std::vectorNewCNode(return_prim, return_inputs); + auto return_prim_c = return_prim->GetPrim(); + MS_ASSERT(return_prim_c != nullptr); + auto return_cnode = anf_graph->NewCNode(return_prim_c, return_inputs); if (return_cnode == nullptr) { MS_LOG(ERROR) << "new cnode error"; return RET_ERROR; @@ -240,7 +241,9 @@ STATUS BuildOpOutputs(const onnx::NodeProto &onnx_node, const FuncGraphPtr &anf_ MS_LOG(ERROR) << "new TupleGetItem failed"; return RET_NULL_PTR; } - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); + auto tuple_get_item_prim_c = tuple_get_item_prim_ptr->GetPrim(); + MS_CHECK_TRUE_MSG(tuple_get_item_prim_c != nullptr, RET_NULL_PTR, "create tuple_get_item_prim_c return nullptr"); + auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_c); MS_CHECK_TRUE_MSG(tuple_get_item_prim != nullptr, RET_NULL_PTR, "create ValueNode return nullptr"); auto get_item_value = NewValueNode(MakeValue(op_idx)); MS_CHECK_TRUE_MSG(get_item_value != nullptr, RET_NULL_PTR, "create ValueNode return nullptr"); @@ -351,7 +354,9 @@ STATUS ConvertGraphOutputs(const onnx::GraphProto &onnx_graph, const FuncGraphPt } make_tuple_inputs.emplace_back(cnode); } - auto make_tuple_cnode = anf_graph->NewCNode(make_tuple_prim_ptr, make_tuple_inputs); + auto make_tuple_prim_c = make_tuple_prim_ptr->GetPrim(); + MS_ASSERT(make_tuple_prim_c != nullptr); + auto make_tuple_cnode = anf_graph->NewCNode(make_tuple_prim_c, make_tuple_inputs); if (make_tuple_cnode == nullptr) { MS_LOG(ERROR) << "new cnode error"; return RET_ERROR; @@ -434,6 +439,11 @@ FuncGraphPtr BuildCondGraph(const AnfNodePtr &root_while_node, int inputs_num, c cond_graph->set_attr("graph_name", MakeValue(cond_graph_name)); return cond_graph; } + +FuncGraphPtr ConvertGraph(api::FuncGraphPtr func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} } // namespace FuncGraphPtr OnnxModelParser::BuildBodyGraph(const onnx::NodeProto &loop_node, const onnx::GraphProto &subgraph_proto, @@ -554,8 +564,9 @@ STATUS CheckOnnxModel(const onnx::GraphProto &onnx_graph) { api::FuncGraphPtr OnnxModelParser::Parse(const converter::ConverterParameters &flag) { auto model_file = flag.model_file; NotSupportOp::GetInstance()->set_fmk_type("ONNX"); - res_graph_ = std::make_shared(); - MS_CHECK_TRUE_MSG(res_graph_ != nullptr, nullptr, "create FuncGraph failed"); + auto graph = std::make_shared(); + MS_CHECK_TRUE_MSG(graph != nullptr, nullptr, "create FuncGraph failed"); + res_graph_ = api::MakeShared(graph); auto status = InitOriginModel(model_file); if (RET_OK != status) { ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); @@ -564,30 +575,28 @@ api::FuncGraphPtr OnnxModelParser::Parse(const converter::ConverterParameters &f } MS_ASSERT(onnx_root_graph_ != nullptr); - auto func_graph = std::dynamic_pointer_cast(res_graph_); - MS_CHECK_TRUE_RET(func_graph != nullptr, nullptr); - status = ConvertOnnxGraph(onnx_root_graph_, func_graph, &anf_nodes_map_, {}, "root_node"); + status = ConvertOnnxGraph(onnx_root_graph_, graph, &anf_nodes_map_, {}, "root_node"); if (RET_OK != status) { ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); MS_LOG(ERROR) << "convert onnx graph failed."; return nullptr; } - static auto root_func_manager = Manage(func_graph); + static auto root_func_manager = Manage(graph); MS_ASSERT(root_func_manager != nullptr); for (auto &subgraph : all_subgraphs_) { MS_ASSERT(subgraph != nullptr); subgraph->set_manager(root_func_manager); subgraph->set_attr("fmk", MakeValue(static_cast(converter::kFmkTypeOnnx))); } - res_graph_->set_attr("graph_name", MakeValue("main_graph")); - res_graph_->set_attr("fmk", MakeValue(static_cast(converter::kFmkTypeOnnx))); - if ((status = CommonAnfAdjust(func_graph)) != RET_OK) { + graph->set_attr("graph_name", MakeValue("main_graph")); + graph->set_attr("fmk", MakeValue(static_cast(converter::kFmkTypeOnnx))); + if ((status = CommonAnfAdjust(graph)) != RET_OK) { MS_LOG(ERROR) << "AdjustForAnf failed."; ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; } std::set all_func_graphs = {}; - GetAllFuncGraph(func_graph, &all_func_graphs); + GetAllFuncGraph(graph, &all_func_graphs); if ((status = Onnx2AnfAdjust(all_func_graphs)) != RET_OK) { MS_LOG(ERROR) << "Onnx2AnfAdjust failed."; ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); @@ -595,7 +604,7 @@ api::FuncGraphPtr OnnxModelParser::Parse(const converter::ConverterParameters &f } auto unify_format = std::make_shared(kFmkTypeOnnx, false); MS_CHECK_TRUE_MSG(unify_format != nullptr, nullptr, "create unify_format return nullptr"); - if (!unify_format->Run(func_graph)) { + if (!unify_format->Run(graph)) { MS_LOG(ERROR) << "Run insert transpose failed."; return nullptr; } @@ -604,6 +613,7 @@ api::FuncGraphPtr OnnxModelParser::Parse(const converter::ConverterParameters &f STATUS OnnxModelParser::InitOriginModel(const std::string &model_file) { MS_ASSERT(res_graph_ != nullptr); + auto res_graph = ConvertGraph(res_graph_); auto status = ValidateFileStr(model_file, ".onnx"); if (status != RET_OK) { MS_LOG(ERROR) << "INPUT ILLEGAL: modelFile must be *.onnx"; @@ -621,7 +631,7 @@ STATUS OnnxModelParser::InitOriginModel(const std::string &model_file) { onnx_root_graph_ = onnx_model_.graph(); auto fmk_value_node = MakeValue(static_cast(converter::kFmkTypeOnnx)); CHECK_NULL_RETURN(fmk_value_node); - res_graph_->set_attr("fmk", fmk_value_node); + res_graph->set_attr("fmk", fmk_value_node); return RET_OK; } @@ -680,10 +690,10 @@ STATUS OnnxModelParser::ConvertNodes(const onnx::GraphProto &onnx_graph, const F MS_ASSERT(anf_graph != nullptr && anf_nodes_map != nullptr && extra_subgraph_inputs != nullptr); STATUS status = RET_OK; for (const auto &onnx_node : onnx_graph.node()) { - ops::PrimitiveC *primitive_c = nullptr; + ops::PrimitiveCPtr primitive_c; auto node_parser = registry::NodeParserRegistry::GetNodeParser(kFmkTypeOnnx, onnx_node.op_type()); if (node_parser != nullptr) { - primitive_c = node_parser->Parse(onnx_graph, onnx_node); + primitive_c = node_parser->Parse(onnx_graph, onnx_node)->GetPrim(); } else { auto node_parser_builtin = OnnxNodeParserRegistry::GetInstance().GetNodeParser(onnx_node.op_type()); if (node_parser_builtin == nullptr) { @@ -853,7 +863,7 @@ STATUS OnnxModelParser::ConvertIfOnnxNode(const onnx::NodeProto &onnx_node, STATUS OnnxModelParser::BuildCNode(const onnx::NodeProto &onnx_node, const FuncGraphPtr &anf_graph, std::unordered_map *anf_nodes_map, - std::vector *graph_inputs, ops::PrimitiveC *primitive_c, + std::vector *graph_inputs, PrimitiveCPtr primitive_c, std::string loop_name) { MS_ASSERT(anf_graph != nullptr && anf_nodes_map != nullptr && graph_inputs != nullptr && primitive_c != nullptr); std::vector op_inputs; @@ -931,7 +941,7 @@ STATUS OnnxModelParser::BuildCNode(const onnx::NodeProto &onnx_node, const FuncG } } } - auto new_cnode = anf_graph->NewCNode(std::shared_ptr(primitive_c), op_inputs); + auto new_cnode = anf_graph->NewCNode(primitive_c, op_inputs); if (new_cnode == nullptr) { MS_LOG(ERROR) << "new cnode error"; return RET_ERROR; @@ -941,7 +951,7 @@ STATUS OnnxModelParser::BuildCNode(const onnx::NodeProto &onnx_node, const FuncG return status; } -STATUS OnnxModelParser::ConvertOpQuantParams(const onnx::NodeProto &onnx_node, ops::PrimitiveC *primitive_c) { +STATUS OnnxModelParser::ConvertOpQuantParams(const onnx::NodeProto &onnx_node, ops::PrimitiveCPtr primitive_c) { MS_ASSERT(primitive_c != nullptr); auto status = ParseQuantParam(onnx_node); if (status != RET_OK) { @@ -1123,7 +1133,9 @@ STATUS OnnxModelParser::AddTensorListStackNode(const AnfNodePtr &root_while_node return RET_ERROR; } tensor_list_stack_prim->set_num_elements(-1); - auto stack_value_node = NewValueNode(tensor_list_stack_prim); + auto prim_c = tensor_list_stack_prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, RET_ERROR); + auto stack_value_node = NewValueNode(prim_c); MS_CHECK_TRUE_MSG(stack_value_node != nullptr, RET_NULL_PTR, "create stack_value_node return nullptr"); std::vector stack_inputs = {stack_value_node, while_output_node, stack_elem_node}; auto tensorlist_stack_cnode = root_anf_graph->NewCNode(stack_inputs); @@ -1288,7 +1300,8 @@ STATUS OnnxModelParser::BuildParameterNodeForQuantParam(const void *data, const MS_LOG(ERROR) << "quant param type don't support."; return RET_NOT_SUPPORT; } - auto parameter_node = res_graph_->add_parameter(); + auto res_graph = ConvertGraph(res_graph_); + auto parameter_node = res_graph->add_parameter(); MS_CHECK_TRUE_MSG(parameter_node != nullptr, RET_NULL_PTR, "create parameter return nullptr"); auto abstract_tensor = CreateTensorAbstract({}, type); if (abstract_tensor == nullptr) { diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.h index f92cdd2adc..02cd8ff82a 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_model_parser.h @@ -54,8 +54,8 @@ class OnnxModelParser : public converter::ModelParser { STATUS BuildParameterNodeForQuantParam(const void *data, const std::string &name, TypeId type); STATUS BuildCNode(const onnx::NodeProto &onnx_node, const FuncGraphPtr &func_graph_ptr, std::unordered_map *anf_nodes_map, std::vector *graph_inputs, - ops::PrimitiveC *primitive_c, std::string loop_name); - STATUS ConvertOpQuantParams(const onnx::NodeProto &onnx_node, ops::PrimitiveC *primitive_c); + ops::PrimitiveCPtr primitive_c, std::string loop_name); + STATUS ConvertOpQuantParams(const onnx::NodeProto &onnx_node, ops::PrimitiveCPtr primitive_c); STATUS ParseQuantParam(const onnx::NodeProto &onnx_node); STATUS SetTensorQuantParam(const std::string &tensor_name, std::vector *quant_params); STATUS SetTensorQuantParamFromNode(const std::string &tensor_name, std::vector *quant_params); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_node_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_node_parser.h index 43e8ba13ac..38fcd78277 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_node_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_node_parser.h @@ -28,6 +28,8 @@ #include "ir/dtype/type_id.h" #include "ops/primitive_c.h" #include "mindspore/core/utils/check_convert_utils.h" +#include "tools/converter/parser/parser_utils.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { @@ -37,9 +39,7 @@ class OnnxNodeParser { virtual ~OnnxNodeParser() = default; - virtual ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { - return nullptr; - } + virtual PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { return nullptr; } static STATUS set_opset_version(int64_t version) { opset_version_ = version; diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.cc index 203113c081..deb77c2313 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.cc @@ -21,8 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxNonMaxSuppressionParser::Parse(const onnx::GraphProto &onnx_graph, - const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxNonMaxSuppressionParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -34,7 +33,7 @@ ops::PrimitiveC *OnnxNonMaxSuppressionParser::Parse(const onnx::GraphProto &onnx } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxNonMaxSuppressionParser("NonMaxSuppression", new OnnxNonMaxSuppressionParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.h index 6c51e67912..780fdfdf1d 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_non_max_suppression_parser.h @@ -27,7 +27,7 @@ class OnnxNonMaxSuppressionParser : public OnnxNodeParser { OnnxNonMaxSuppressionParser() : OnnxNodeParser("NonMaxSuppression") {} ~OnnxNonMaxSuppressionParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.cc index b4b0c862b8..daa973c4e6 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.cc @@ -22,10 +22,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxNonZeroParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxNonZeroParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxNonZeroParser("NonZero", new OnnxNonZeroParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.h index a8ed83448f..2ed9155490 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_nonzero_parser.h @@ -27,7 +27,7 @@ class OnnxNonZeroParser : public OnnxNodeParser { OnnxNonZeroParser() : OnnxNodeParser("NonZero") {} ~OnnxNonZeroParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.cc index a8aaca5a1f..2eb882a1ee 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxOneHotParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxOneHotParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -31,7 +31,7 @@ ops::PrimitiveC *OnnxOneHotParser::Parse(const onnx::GraphProto &onnx_graph, con } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxOneHotParser("OneHot", new OnnxOneHotParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.h index 9ed0a6278b..d39c2c2342 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_onehot_parser.h @@ -27,7 +27,7 @@ class OnnxOneHotParser : public OnnxNodeParser { OnnxOneHotParser() : OnnxNodeParser("OneHot") {} ~OnnxOneHotParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_pad_adjust.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_pad_adjust.cc index 94023fbfc0..0da814866d 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_pad_adjust.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_pad_adjust.cc @@ -36,8 +36,9 @@ CNodePtr NewReshapeOpNode(const FuncGraphPtr &func_graph, const AnfNodePtr &inpu MS_LOG(ERROR) << "create reshape failed."; return nullptr; } - reshape_prim->set_attr("shape", MakeValue(shape)); - ValueNodePtr value_node = NewValueNode(reshape_prim); + auto prim_c = reshape_prim->GetPrim(); + prim_c->set_attr("shape", MakeValue(shape)); + ValueNodePtr value_node = NewValueNode(prim_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "create valuenode return nullptr"); auto new_parameter = opt::BuildIntVecParameterNode(func_graph, shape, input_node->fullname_with_scope() + "_reshape/shape"); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.cc index 0003f237a2..da98a1488c 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.cc @@ -24,9 +24,11 @@ namespace mindspore { namespace lite { constexpr auto kNamePadContiguous = "pad_contiguous"; -ops::PrimitiveC *OnnxPadParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxPadParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); mindspore::PaddingMode padding_mode; for (const auto &onnx_node_attr : onnx_node.attribute()) { const auto &attribute_name = onnx_node_attr.name(); @@ -57,9 +59,9 @@ ops::PrimitiveC *OnnxPadParser::Parse(const onnx::GraphProto &onnx_graph, const prim->set_constant_value(onnx_node_attr.f()); } } - prim->AddAttr(kNamePadContiguous, MakeValue(true)); + prim_c->AddAttr(kNamePadContiguous, MakeValue(true)); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxPadParser("Pad", new OnnxPadParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.h index 641b35f39a..020e0a765f 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_pad_parser.h @@ -27,7 +27,7 @@ class OnnxPadParser : public OnnxNodeParser { OnnxPadParser() : OnnxNodeParser("Pad") {} ~OnnxPadParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.cc index 59c203d774..6a3d0a30f9 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.cc @@ -24,10 +24,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxAvgPoolParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxAvgPoolParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); prim->set_pad_mode(mindspore::PadMode::PAD); mindspore::RoundMode round_mode = mindspore::RoundMode::FLOOR; std::vector kernels; @@ -94,14 +96,16 @@ ops::PrimitiveC *OnnxAvgPoolParser::Parse(const onnx::GraphProto &onnx_graph, co } int fmk_type = converter::FmkType::kFmkTypeOnnx; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim_c->AddAttr(ops::kFmkType, MakeValue(fmk_type)); + return prim->GetPrim(); } -ops::PrimitiveC *OnnxMaxPoolParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxMaxPoolParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); mindspore::RoundMode round_mode = mindspore::RoundMode::FLOOR; std::vector kernels; std::vector strides; @@ -166,8 +170,8 @@ ops::PrimitiveC *OnnxMaxPoolParser::Parse(const onnx::GraphProto &onnx_graph, co prim->set_global(onnx_node.op_type() == "GlobalMaxPool"); int fmk_type = converter::FmkType::kFmkTypeOnnx; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + prim_c->AddAttr(ops::kFmkType, MakeValue(fmk_type)); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxAveragePoolParser("AveragePool", new OnnxAvgPoolParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.h index 0fc82ba857..075773637d 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_pool_parser.h @@ -27,7 +27,7 @@ class OnnxAvgPoolParser : public OnnxNodeParser { OnnxAvgPoolParser() : OnnxNodeParser("AvgPool") {} ~OnnxAvgPoolParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; class OnnxMaxPoolParser : public OnnxNodeParser { @@ -35,7 +35,7 @@ class OnnxMaxPoolParser : public OnnxNodeParser { OnnxMaxPoolParser() : OnnxNodeParser("MaxPool") {} ~OnnxMaxPoolParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.cc index a7a223549e..6a14d0617f 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxQuantizeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxQuantizeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); if (onnx_node.op_type() == "Int8Quantize") { @@ -35,7 +35,7 @@ ops::PrimitiveC *OnnxQuantizeParser::Parse(const onnx::GraphProto &onnx_graph, c return nullptr; } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxInt8QuantizeParser("Int8Quantize", new OnnxQuantizeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.h index 0b6cbc2898..f5ac6351af 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_quantize_parser.h @@ -27,7 +27,7 @@ class OnnxQuantizeParser : public OnnxNodeParser { OnnxQuantizeParser() : OnnxNodeParser("Quantize") {} ~OnnxQuantizeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.cc index 22267e9336..72c67f14cd 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.cc @@ -22,19 +22,19 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxRandomNormalParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxRandomNormalParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { if (onnx_node_attr.name() == "dtype") { auto onnx_dtype = static_cast(onnx_node_attr.i()); auto data_type = OnnxNodeParser::GetDataTypeFromOnnx(onnx_dtype); - prim->AddAttr(ops::kDataType, MakeValue(static_cast(data_type))); + prim->AddAttr(ops::kDataType, api::MakeValue(static_cast(data_type))); } else if (onnx_node_attr.name() == "shape") { std::vector shape; std::transform(onnx_node_attr.ints().begin(), onnx_node_attr.ints().end(), std::back_inserter(shape), [](int ele) { return static_cast(ele); }); - prim->AddAttr(ops::kShape, MakeValue(shape)); + prim->AddAttr(ops::kShape, api::MakeValue(shape)); } else if (onnx_node_attr.name() == "seed") { prim->set_seed(static_cast(onnx_node_attr.f())); } else if (onnx_node_attr.name() == "mean") { @@ -43,7 +43,7 @@ ops::PrimitiveC *OnnxRandomNormalParser::Parse(const onnx::GraphProto &onnx_grap prim->set_scale(static_cast(onnx_node_attr.f())); } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxRandomNormalParser("RandomNormal", new OnnxRandomNormalParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.h index 655274760d..8f24c54f48 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_random_normal_parser.h @@ -25,7 +25,7 @@ class OnnxRandomNormalParser : public OnnxNodeParser { OnnxRandomNormalParser() : OnnxNodeParser("RandomNormal") {} ~OnnxRandomNormalParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.cc index dfa98cd59c..1635f0b173 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.cc @@ -21,12 +21,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxRangeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxRangeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_d_type(0); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxRangeParser("Range", new OnnxRangeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.h index 22f8ffecaa..5c71af8d8e 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_range_parser.h @@ -27,7 +27,7 @@ class OnnxRangeParser : public OnnxNodeParser { OnnxRangeParser() : OnnxNodeParser("Range") {} ~OnnxRangeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.cc index a972e5ab72..0a82fe93b5 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.cc @@ -22,11 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxReduceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxReduceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_keep_dims(true); - + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); std::vector axes = {}; for (const auto &onnx_node_attr : onnx_node.attribute()) { const auto &attribute_name = onnx_node_attr.name(); @@ -35,7 +36,7 @@ ops::PrimitiveC *OnnxReduceParser::Parse(const onnx::GraphProto &onnx_graph, con for (int i = 0; i < size; ++i) { axes.push_back(onnx_node_attr.ints(i)); } - prim->AddAttr("axes", MakeValue(axes)); + prim_c->AddAttr("axes", MakeValue(axes)); } else if (attribute_name == "keepdims") { prim->set_keep_dims(static_cast(onnx_node_attr.i())); } @@ -59,7 +60,7 @@ ops::PrimitiveC *OnnxReduceParser::Parse(const onnx::GraphProto &onnx_graph, con return nullptr; } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxReduceMeanParser("ReduceMean", new OnnxReduceParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.h index 95080675d8..b0527c8b99 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_reduce_parser.h @@ -27,7 +27,7 @@ class OnnxReduceParser : public OnnxNodeParser { OnnxReduceParser() : OnnxNodeParser("Reduce") {} ~OnnxReduceParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.cc index 57d790a1c5..601f198ff1 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.cc @@ -22,9 +22,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxReshapeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxReshapeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); std::vector shape; shape.clear(); if (onnx_node.input_size() != 2) { @@ -34,12 +36,12 @@ ops::PrimitiveC *OnnxReshapeParser::Parse(const onnx::GraphProto &onnx_graph, co for (int i = 0; i < onnx_node_attr.ints_size(); ++i) { shape.push_back(static_cast(onnx_node_attr.ints(i))); } - prim->AddAttr("shape", MakeValue(shape)); + prim_c->AddAttr("shape", MakeValue(shape)); } } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxReshapeParser("Reshape", new OnnxReshapeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.h index cc5a252574..0841b261b1 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_reshape_parser.h @@ -27,7 +27,7 @@ class OnnxReshapeParser : public OnnxNodeParser { OnnxReshapeParser() : OnnxNodeParser("Reshape") {} ~OnnxReshapeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.cc index 562109b28f..6c89a7f5a7 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.cc @@ -25,10 +25,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxResizeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxResizeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NCHW)); prim->set_nearest_mode(mindspore::NearestMode::ROUND_HALF_DOWN); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -79,7 +81,7 @@ ops::PrimitiveC *OnnxResizeParser::Parse(const onnx::GraphProto &onnx_graph, con } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxResizeParser("Resize", new OnnxResizeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.h index fc19617032..7586d9d69c 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_resize_parser.h @@ -27,7 +27,7 @@ class OnnxResizeParser : public OnnxNodeParser { OnnxResizeParser() : OnnxNodeParser("Resize") {} ~OnnxResizeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.cc index 9b34d10d76..a5e717e035 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.cc @@ -21,10 +21,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxScatterNdParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxScatterNdParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxScatterNdParser("ScatterND", new OnnxScatterNdParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.h index d21f67e476..656fde3111 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_scatter_nd_parser.h @@ -27,7 +27,7 @@ class OnnxScatterNdParser : public OnnxNodeParser { OnnxScatterNdParser() : OnnxNodeParser("ScatterNd") {} ~OnnxScatterNdParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.cc index da4dd6ac55..1d718b54a9 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.cc @@ -21,10 +21,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxShapeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxShapeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxShapeParser("Shape", new OnnxShapeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.h index b78d4673b1..ce463ff0ee 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_shape_parser.h @@ -27,7 +27,7 @@ class OnnxShapeParser : public OnnxNodeParser { OnnxShapeParser() : OnnxNodeParser("Shape") {} ~OnnxShapeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.cc index c54a75c6ec..cf439a8361 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.cc @@ -28,9 +28,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxSliceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSliceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); std::vector starts; std::vector ends; std::vector axes; @@ -75,7 +77,7 @@ ops::PrimitiveC *OnnxSliceParser::Parse(const onnx::GraphProto &onnx_graph, cons size = static_cast(steps.size()); } if (size == -1) { - return prim.release(); + return prim->GetPrim(); } if (axes.empty()) { for (size_t i = 0; i < starts.size(); ++i) { @@ -86,14 +88,14 @@ ops::PrimitiveC *OnnxSliceParser::Parse(const onnx::GraphProto &onnx_graph, cons steps.assign(starts.size(), 1); } - prim->AddAttr("starts", MakeValue(starts)); - prim->AddAttr("axes", MakeValue(axes)); - prim->AddAttr("ends", MakeValue(ends)); - prim->AddAttr("steps", MakeValue(steps)); - int fmk_type = converter::FmkType::kFmkTypeOnnx; - prim->AddAttr(ops::kFmkType, MakeValue(fmk_type)); + prim_c->AddAttr("starts", MakeValue(starts)); + prim_c->AddAttr("axes", MakeValue(axes)); + prim_c->AddAttr("ends", MakeValue(ends)); + prim_c->AddAttr("steps", MakeValue(steps)); + int64_t fmk_type = converter::FmkType::kFmkTypeOnnx; + prim_c->AddAttr(ops::kFmkType, MakeValue(fmk_type)); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxSliceParser("Slice", new OnnxSliceParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.h index 9608931e9f..9fce64fa3f 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_slice_parser.h @@ -29,7 +29,7 @@ class OnnxSliceParser : public OnnxNodeParser { OnnxSliceParser() : OnnxNodeParser("Slice") {} ~OnnxSliceParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.cc index f04edd8001..2c2bc0e77e 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxSoftMaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSoftMaxParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); int64_t axis; @@ -38,7 +38,7 @@ ops::PrimitiveC *OnnxSoftMaxParser::Parse(const onnx::GraphProto &onnx_graph, co } prim->set_axis({axis}); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxSoftMaxParser("Softmax", new OnnxSoftMaxParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.h index ccbf24a341..317223f3b4 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_softmax_parser.h @@ -27,7 +27,7 @@ class OnnxSoftMaxParser : public OnnxNodeParser { OnnxSoftMaxParser() : OnnxNodeParser("Softmax") {} ~OnnxSoftMaxParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.cc index 9b2e8c6022..f634d0ec26 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.cc @@ -21,7 +21,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxSpaceToDepthParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSpaceToDepthParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { @@ -31,7 +31,7 @@ ops::PrimitiveC *OnnxSpaceToDepthParser::Parse(const onnx::GraphProto &onnx_grap } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxSpaceToDepthParser("SpaceToDepth", new OnnxSpaceToDepthParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.h index 00798fcfee..b8a813bb67 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_space_to_depth_parser.h @@ -27,7 +27,7 @@ class OnnxSpaceToDepthParser : public OnnxNodeParser { OnnxSpaceToDepthParser() : OnnxNodeParser("SpaceToDepth") {} ~OnnxSpaceToDepthParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.cc index 8fa33043ae..96ead6ada2 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.cc @@ -23,7 +23,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxSpliceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSpliceParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { MS_LOG(DEBUG) << "onnx Splice Parser"; auto primitive = std::make_unique(); MS_CHECK_TRUE_RET(primitive != nullptr, nullptr); @@ -52,7 +52,7 @@ ops::PrimitiveC *OnnxSpliceParser::Parse(const onnx::GraphProto &onnx_graph, con } } primitive->Init(context, forward_indexes, output_dim); - return primitive.release(); + return primitive->GetPrim(); } OnnxNodeRegistrar g_onnxSpliceParser("Splice", new OnnxSpliceParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.h index 07542837c0..4dcc274dd8 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_splice_parser.h @@ -27,7 +27,7 @@ class OnnxSpliceParser : public OnnxNodeParser { OnnxSpliceParser() : OnnxNodeParser("Splice") {} ~OnnxSpliceParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.cc index 930913b1ec..465e4b1830 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.cc @@ -23,7 +23,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxSplitParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSplitParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_axis(0); @@ -45,7 +45,7 @@ ops::PrimitiveC *OnnxSplitParser::Parse(const onnx::GraphProto &onnx_graph, cons } prim->set_output_num(split_num); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxSplitParser("Split", new OnnxSplitParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.h index d2d28e529c..f5405d9ff0 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_split_parser.h @@ -27,7 +27,7 @@ class OnnxSplitParser : public OnnxNodeParser { OnnxSplitParser() : OnnxNodeParser("Split") {} ~OnnxSplitParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.cc index 395e255222..331a42276d 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxSqueezeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxSqueezeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); std::vector axis; @@ -36,7 +36,7 @@ ops::PrimitiveC *OnnxSqueezeParser::Parse(const onnx::GraphProto &onnx_graph, co } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxSqueezeParser("Squeeze", new OnnxSqueezeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.h index a53c35e0b6..b4b54d6893 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_squeeze_parser.h @@ -27,7 +27,7 @@ class OnnxSqueezeParser : public OnnxNodeParser { OnnxSqueezeParser() : OnnxNodeParser("Squeeze") {} ~OnnxSqueezeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.cc index 70974740d4..945d1d1a8a 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.cc @@ -22,10 +22,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxTileParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxTileParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxTileParser("Tile", new OnnxTileParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.h index d03d4e290f..ec21a59a60 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_tile_parser.h @@ -27,7 +27,7 @@ class OnnxTileParser : public OnnxNodeParser { OnnxTileParser() : OnnxNodeParser("Tile") {} ~OnnxTileParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.cc index ed96ae1151..ef87944559 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.cc @@ -22,19 +22,21 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxTopkParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxTopkParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); for (const auto &onnx_node_attr : onnx_node.attribute()) { const auto &attribute_name = onnx_node_attr.name(); if (attribute_name == "k") { auto k_value = MakeValue(static_cast(onnx_node_attr.i())); MS_CHECK_TRUE_MSG(k_value != nullptr, nullptr, "CreateValueNode failed"); - prim->AddAttr("k", k_value); + prim_c->AddAttr("k", k_value); } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxTopkParser("TopK", new OnnxTopkParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.h index f075fa42f4..89304c3b06 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_topk_parser.h @@ -27,7 +27,7 @@ class OnnxTopkParser : public OnnxNodeParser { OnnxTopkParser() : OnnxNodeParser("TopK") {} ~OnnxTopkParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.cc index 2a43adf809..74334b8f50 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.cc @@ -22,9 +22,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxTransposeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxTransposeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); std::vector perm; for (const auto &onnx_node_attr : onnx_node.attribute()) { const auto &attribute_name = onnx_node_attr.name(); @@ -37,8 +39,8 @@ ops::PrimitiveC *OnnxTransposeParser::Parse(const onnx::GraphProto &onnx_graph, } auto perm_value = MakeValue(perm); MS_CHECK_TRUE_MSG(perm_value != nullptr, nullptr, "MakeValue failed"); - prim->AddAttr("perm", perm_value); - return prim.release(); + prim_c->AddAttr("perm", perm_value); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxTransposeParser("Transpose", new OnnxTransposeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.h index 4f6f76edce..eb9744c8e4 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_transpose_parser.h @@ -27,7 +27,7 @@ class OnnxTransposeParser : public OnnxNodeParser { OnnxTransposeParser() : OnnxNodeParser("Transpose") {} ~OnnxTransposeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.cc index a98545de15..c71edde8de 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.cc @@ -22,7 +22,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxUnSqueezeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxUnSqueezeParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); std::vector axis; @@ -36,7 +36,7 @@ ops::PrimitiveC *OnnxUnSqueezeParser::Parse(const onnx::GraphProto &onnx_graph, } } - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxUnsqueezeParser("Unsqueeze", new OnnxUnSqueezeParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.h index eb8074e2b4..a76900e349 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_unsqueeze_parser.h @@ -27,7 +27,7 @@ class OnnxUnSqueezeParser : public OnnxNodeParser { OnnxUnSqueezeParser() : OnnxNodeParser("Unsqueeze") {} ~OnnxUnSqueezeParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.cc index 54005554f9..e9d2f96505 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.cc @@ -23,7 +23,7 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxUpsampleParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxUpsampleParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_method(mindspore::ResizeMethod::NEAREST); // use bilinear method @@ -41,7 +41,7 @@ ops::PrimitiveC *OnnxUpsampleParser::Parse(const onnx::GraphProto &onnx_graph, c } prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::ASYMMETRIC); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxUpsampleParser("Upsample", new OnnxUpsampleParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.h index 56ce858faf..b8b99986d9 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_upsample_parser.h @@ -27,7 +27,7 @@ class OnnxUpsampleParser : public OnnxNodeParser { OnnxUpsampleParser() : OnnxNodeParser("Upsample") {} ~OnnxUpsampleParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.cc b/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.cc index 694310ca99..a83ea43f17 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.cc +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.cc @@ -21,10 +21,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *OnnxWhereParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { +PrimitiveCPtr OnnxWhereParser::Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } OnnxNodeRegistrar g_onnxWhereParser("Where", new OnnxWhereParser()); diff --git a/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.h b/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.h index 9f09f193dc..f73901e1a7 100644 --- a/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.h +++ b/mindspore/lite/tools/converter/parser/onnx/onnx_where_parser.h @@ -27,7 +27,7 @@ class OnnxWhereParser : public OnnxNodeParser { OnnxWhereParser() : OnnxNodeParser("Where") {} ~OnnxWhereParser() override = default; - ops::PrimitiveC *Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; + PrimitiveCPtr Parse(const onnx::GraphProto &onnx_graph, const onnx::NodeProto &onnx_node) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/parser_utils.cc b/mindspore/lite/tools/converter/parser/parser_utils.cc index 92f735a666..8d12156a07 100644 --- a/mindspore/lite/tools/converter/parser/parser_utils.cc +++ b/mindspore/lite/tools/converter/parser/parser_utils.cc @@ -13,6 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/parser/parser_utils.h" #include #include @@ -24,6 +25,7 @@ #include "ops/apply_momentum.h" #include "ops/fusion/conv2d_fusion.h" #include "ops/fusion/conv2d_transpose_fusion.h" +#include "ops/fusion/conv2d_backprop_input_fusion.h" #include "ops/sgd.h" #include "tools/common/tensor_util.h" #include "tools/converter/parser/conv1d_inout_adjust.h" @@ -34,6 +36,7 @@ #include "tools/optimizer/common/gllo_utils.h" #include "tools/optimizer/format/to_format_base.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::lite { namespace { diff --git a/mindspore/lite/tools/converter/parser/parser_utils.h b/mindspore/lite/tools/converter/parser/parser_utils.h index 25caaae0e2..4145976a80 100644 --- a/mindspore/lite/tools/converter/parser/parser_utils.h +++ b/mindspore/lite/tools/converter/parser/parser_utils.h @@ -19,6 +19,8 @@ #include #include +#include +#include "ops/primitive_c.h" #include "include/registry/model_parser.h" #include "ir/anf.h" #include "ir/func_graph.h" @@ -27,6 +29,7 @@ namespace mindspore { namespace lite { +using PrimitiveCPtr = std::shared_ptr; void GetAllFuncGraph(const FuncGraphPtr &func_graph, std::set *all_func_graphs); int CommonAnfAdjust(const FuncGraphPtr &func_graph); int GetTransposePerm(schema::Format src_format, schema::Format dst_format, std::vector *perm); diff --git a/mindspore/lite/tools/converter/parser/tf/functionalize_cond.cc b/mindspore/lite/tools/converter/parser/tf/functionalize_cond.cc index 7ad5252f13..08f6cb7efd 100644 --- a/mindspore/lite/tools/converter/parser/tf/functionalize_cond.cc +++ b/mindspore/lite/tools/converter/parser/tf/functionalize_cond.cc @@ -163,7 +163,8 @@ FuncGraphPtr FunctionalizeCond::CreateBranchGraph(const AnfNodePtr &node, std::s MS_LOG(ERROR) << "GetReturnPrim return nullptr"; return nullptr; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + auto value_node = NewValueNode(return_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); std::vector op_inputs{value_node, node}; // If subgraph only has one output tensor auto return_cnode = graph->NewCNode(op_inputs); diff --git a/mindspore/lite/tools/converter/parser/tf/functionalize_cond.h b/mindspore/lite/tools/converter/parser/tf/functionalize_cond.h index ef1e461505..ca05bc6324 100644 --- a/mindspore/lite/tools/converter/parser/tf/functionalize_cond.h +++ b/mindspore/lite/tools/converter/parser/tf/functionalize_cond.h @@ -15,6 +15,7 @@ */ #ifndef MINDSPORE_LITE_TOOLS_OPTIMIZER_GRAPH_FUNCTIONALIZE_COND_H_ #define MINDSPORE_LITE_TOOLS_OPTIMIZER_GRAPH_FUNCTIONALIZE_COND_H_ +#define USE_DEPRECATED_API #include #include diff --git a/mindspore/lite/tools/converter/parser/tf/functionalize_while.cc b/mindspore/lite/tools/converter/parser/tf/functionalize_while.cc index e8fa4dca00..e50ae835dd 100644 --- a/mindspore/lite/tools/converter/parser/tf/functionalize_while.cc +++ b/mindspore/lite/tools/converter/parser/tf/functionalize_while.cc @@ -234,7 +234,9 @@ STATUS FunctionalizeWhile::UpdateExitNodeUser() { MS_LOG(ERROR) << "GetTupleGetItemPrim return nullptr"; return RET_NULL_PTR; } - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); + auto tuple_get_item_prim_c = tuple_get_item_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(tuple_get_item_prim_c); + auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_c); CHECK_NULL_RETURN(tuple_get_item_prim); const auto &exit_node = node; auto switch_node = BlongToWhichSwitch(exit_node); @@ -380,7 +382,9 @@ STATUS FunctionalizeWhile::IdentifyCondSubgraphOutput() { MS_LOG(ERROR) << "GetReturnPrim return nullptr"; return RET_NULL_PTR; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(return_prim_c); + auto value_node = NewValueNode(return_prim_c); if (value_node == nullptr) { MS_LOG(ERROR) << "new value_node failed."; return RET_NULL_PTR; @@ -541,7 +545,9 @@ STATUS FunctionalizeWhile::IdentifyBodySubgraphOutput() { MS_LOG(ERROR) << "GetReturnPrim return nullptr"; return RET_NULL_PTR; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(return_prim_c); + auto value_node = NewValueNode(return_prim_c); CHECK_NULL_RETURN(value_node); // cond subgraph output is LoopCond's input std::vector op_inputs{value_node}; @@ -558,7 +564,9 @@ STATUS FunctionalizeWhile::IdentifyBodySubgraphOutput() { MS_LOG(ERROR) << "GetMakeTuplePrim return nullptr"; return RET_NULL_PTR; } - auto make_tuple_prim = NewValueNode(make_tuple_prim_ptr); + auto prim_c = make_tuple_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(prim_c); + auto make_tuple_prim = NewValueNode(prim_c); CHECK_NULL_RETURN(make_tuple_prim); make_tuple_inputs.insert(make_tuple_inputs.begin(), make_tuple_prim); auto make_tuple_cnode = body_sub_func_graph_->NewCNode(make_tuple_inputs); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.cc index 4bb96c0c80..00b7675da9 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.cc @@ -24,9 +24,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFActivationParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFActivationParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); if (tf_op.op() == "Relu") { @@ -64,7 +64,7 @@ ops::PrimitiveC *TFActivationParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfReluParser("Relu", new TFActivationParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.h index ece8aee5fd..53cd08c998 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_activation_parser.h @@ -29,9 +29,9 @@ class TFActivationParser : public TFNodeParser { TFActivationParser() = default; ~TFActivationParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.cc index ee25910973..54a67d5581 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFArgMaxParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFArgMaxParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -50,7 +50,7 @@ ops::PrimitiveC *TFArgMaxParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfArgMaxParser("ArgMax", new TFArgMaxParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.h index 197afb60ba..b42c6df322 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_argmax_parser.h @@ -28,9 +28,9 @@ class TFArgMaxParser : public TFNodeParser { TFArgMaxParser() = default; ~TFArgMaxParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.cc index c41aa32aaf..5fb8e03809 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFArgMinParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFArgMinParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -51,7 +51,7 @@ ops::PrimitiveC *TFArgMinParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfArgMinParser("ArgMin", new TFArgMinParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.h index 1827eee910..b0c01c4c33 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_argmin_parser.h @@ -28,9 +28,9 @@ class TFArgMinParser : public TFNodeParser { TFArgMinParser() = default; ~TFArgMinParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.cc index 364e9d2d59..5351770fc2 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.cc @@ -49,9 +49,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFAddParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFAddParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -60,12 +60,12 @@ ops::PrimitiveC *TFAddParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFSubParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSubParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -74,12 +74,12 @@ ops::PrimitiveC *TFSubParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFMulParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFMulParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -88,27 +88,29 @@ ops::PrimitiveC *TFMulParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFDivParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFDivParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); *output_size = 1; if (AddOpInput(tf_op, 0, inputs) != RET_OK || AddOpInput(tf_op, 1, inputs) != RET_OK) { MS_LOG(ERROR) << "add op input failed"; return nullptr; } std::string original_name = tf_op.op(); - prim->AddAttr(ops::kOriginalOpName, MakeValue(original_name)); - return prim.release(); + prim_c->AddAttr(ops::kOriginalOpName, MakeValue(original_name)); + return prim->GetPrim(); } -ops::PrimitiveC *TFMaximumParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFMaximumParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -117,12 +119,12 @@ ops::PrimitiveC *TFMaximumParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFMinimumParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFMinimumParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -131,12 +133,12 @@ ops::PrimitiveC *TFMinimumParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFGreaterParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFGreaterParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -145,12 +147,12 @@ ops::PrimitiveC *TFGreaterParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFGreaterEqualParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFGreaterEqualParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -159,12 +161,12 @@ ops::PrimitiveC *TFGreaterEqualParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFLessParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLessParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -173,12 +175,12 @@ ops::PrimitiveC *TFLessParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFLessEqualParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLessEqualParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -187,12 +189,12 @@ ops::PrimitiveC *TFLessEqualParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFEqualParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFEqualParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -201,12 +203,12 @@ ops::PrimitiveC *TFEqualParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFNotEqualParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFNotEqualParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -215,12 +217,12 @@ ops::PrimitiveC *TFNotEqualParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFSquaredDifferenceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSquaredDifferenceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -229,12 +231,12 @@ ops::PrimitiveC *TFSquaredDifferenceParser::Parse(const tensorflow::NodeDef &tf_ return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFRsqrtParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFRsqrtParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -242,12 +244,12 @@ ops::PrimitiveC *TFRsqrtParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFRoundParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFRoundParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -256,12 +258,12 @@ ops::PrimitiveC *TFRoundParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFCeilParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFCeilParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -270,12 +272,12 @@ ops::PrimitiveC *TFCeilParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFExpParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFExpParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -284,12 +286,12 @@ ops::PrimitiveC *TFExpParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFFloorParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFFloorParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -298,12 +300,12 @@ ops::PrimitiveC *TFFloorParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFFloorDivParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFFloorDivParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -312,12 +314,12 @@ ops::PrimitiveC *TFFloorDivParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFFloorModParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFFloorModParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -326,12 +328,12 @@ ops::PrimitiveC *TFFloorModParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFLogParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLogParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -340,12 +342,12 @@ ops::PrimitiveC *TFLogParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFSqrtParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSqrtParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -354,12 +356,12 @@ ops::PrimitiveC *TFSqrtParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFCosParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFCosParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -368,12 +370,12 @@ ops::PrimitiveC *TFCosParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFSinParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSinParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -382,12 +384,12 @@ ops::PrimitiveC *TFSinParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFSquareParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSquareParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -396,12 +398,12 @@ ops::PrimitiveC *TFSquareParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFPowParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFPowParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -410,12 +412,12 @@ ops::PrimitiveC *TFPowParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFAbsParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFAbsParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -424,7 +426,7 @@ ops::PrimitiveC *TFAbsParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfAddParser("Add", new TFAddParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.h index 98cf5551d4..87156dc970 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_arithmetic_parser.h @@ -29,9 +29,9 @@ class TFAddParser : public TFNodeParser { TFAddParser() = default; ~TFAddParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFSubParser : public TFNodeParser { @@ -39,9 +39,9 @@ class TFSubParser : public TFNodeParser { TFSubParser() = default; ~TFSubParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFMulParser : public TFNodeParser { @@ -49,9 +49,9 @@ class TFMulParser : public TFNodeParser { TFMulParser() = default; ~TFMulParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFDivParser : public TFNodeParser { @@ -59,9 +59,9 @@ class TFDivParser : public TFNodeParser { TFDivParser() = default; ~TFDivParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFMaximumParser : public TFNodeParser { @@ -69,9 +69,9 @@ class TFMaximumParser : public TFNodeParser { TFMaximumParser() = default; ~TFMaximumParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFMinimumParser : public TFNodeParser { @@ -79,9 +79,9 @@ class TFMinimumParser : public TFNodeParser { TFMinimumParser() = default; ~TFMinimumParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFGreaterParser : public TFNodeParser { @@ -89,9 +89,9 @@ class TFGreaterParser : public TFNodeParser { TFGreaterParser() = default; ~TFGreaterParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFGreaterEqualParser : public TFNodeParser { @@ -99,9 +99,9 @@ class TFGreaterEqualParser : public TFNodeParser { TFGreaterEqualParser() = default; ~TFGreaterEqualParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFLessParser : public TFNodeParser { @@ -109,9 +109,9 @@ class TFLessParser : public TFNodeParser { TFLessParser() = default; ~TFLessParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFLessEqualParser : public TFNodeParser { @@ -119,9 +119,9 @@ class TFLessEqualParser : public TFNodeParser { TFLessEqualParser() = default; ~TFLessEqualParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFEqualParser : public TFNodeParser { @@ -129,9 +129,9 @@ class TFEqualParser : public TFNodeParser { TFEqualParser() = default; ~TFEqualParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFNotEqualParser : public TFNodeParser { @@ -139,9 +139,9 @@ class TFNotEqualParser : public TFNodeParser { TFNotEqualParser() = default; ~TFNotEqualParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFSquaredDifferenceParser : public TFNodeParser { @@ -149,9 +149,9 @@ class TFSquaredDifferenceParser : public TFNodeParser { TFSquaredDifferenceParser() = default; ~TFSquaredDifferenceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFRsqrtParser : public TFNodeParser { @@ -159,9 +159,9 @@ class TFRsqrtParser : public TFNodeParser { TFRsqrtParser() = default; ~TFRsqrtParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFRoundParser : public TFNodeParser { @@ -169,9 +169,9 @@ class TFRoundParser : public TFNodeParser { TFRoundParser() = default; ~TFRoundParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFCeilParser : public TFNodeParser { @@ -179,9 +179,9 @@ class TFCeilParser : public TFNodeParser { TFCeilParser() = default; ~TFCeilParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFExpParser : public TFNodeParser { @@ -189,9 +189,9 @@ class TFExpParser : public TFNodeParser { TFExpParser() = default; ~TFExpParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFFloorParser : public TFNodeParser { @@ -199,9 +199,9 @@ class TFFloorParser : public TFNodeParser { TFFloorParser() = default; ~TFFloorParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFFloorDivParser : public TFNodeParser { @@ -209,9 +209,9 @@ class TFFloorDivParser : public TFNodeParser { TFFloorDivParser() = default; ~TFFloorDivParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFFloorModParser : public TFNodeParser { @@ -219,9 +219,9 @@ class TFFloorModParser : public TFNodeParser { TFFloorModParser() = default; ~TFFloorModParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFLogParser : public TFNodeParser { @@ -229,9 +229,9 @@ class TFLogParser : public TFNodeParser { TFLogParser() = default; ~TFLogParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFSqrtParser : public TFNodeParser { @@ -239,9 +239,9 @@ class TFSqrtParser : public TFNodeParser { TFSqrtParser() = default; ~TFSqrtParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFCosParser : public TFNodeParser { @@ -249,9 +249,9 @@ class TFCosParser : public TFNodeParser { TFCosParser() = default; ~TFCosParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFSinParser : public TFNodeParser { @@ -259,9 +259,9 @@ class TFSinParser : public TFNodeParser { TFSinParser() = default; ~TFSinParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFSquareParser : public TFNodeParser { @@ -269,9 +269,9 @@ class TFSquareParser : public TFNodeParser { TFSquareParser() = default; ~TFSquareParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFPowParser : public TFNodeParser { @@ -279,9 +279,9 @@ class TFPowParser : public TFNodeParser { TFPowParser() = default; ~TFPowParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFAbsParser : public TFNodeParser { @@ -289,9 +289,9 @@ class TFAbsParser : public TFNodeParser { TFAbsParser() = default; ~TFAbsParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.cc index d228868b5a..75648e758f 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFAssertParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFAssertParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -43,7 +43,7 @@ ops::PrimitiveC *TFAssertParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfAssertParser("Assert", new TFAssertParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.h index b1f3b1cc52..55a06b5911 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_assert_parser.h @@ -28,9 +28,9 @@ class TFAssertParser : public TFNodeParser { TFAssertParser() = default; ~TFAssertParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.cc index 270ba37ebd..3b51eb3619 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFBatchMatMulParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFBatchMatMulParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -44,7 +44,7 @@ ops::PrimitiveC *TFBatchMatMulParser::Parse(const tensorflow::NodeDef &tf_op, for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfBatchMatMulParser("BatchMatMul", new TFBatchMatMulParser()); TFNodeRegistrar g_tfBatchMatMulV2Parser("BatchMatMulV2", new TFBatchMatMulParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.h index 287e295700..66d2cef632 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_batch_matmul_parser.h @@ -28,9 +28,9 @@ class TFBatchMatMulParser : public TFNodeParser { TFBatchMatMulParser() = default; ~TFBatchMatMulParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.cc index e26392af8b..f68bdece66 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFBatchToSpaceNDParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFBatchToSpaceNDParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -36,7 +36,7 @@ ops::PrimitiveC *TFBatchToSpaceNDParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfBatchToSpaceNDParser("BatchToSpaceND", new TFBatchToSpaceNDParser()); TFNodeRegistrar g_tfBatchToSpaceParser("BatchToSpace", new TFBatchToSpaceNDParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.h index 3c5befaa57..45527afdfe 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_batch_to_space_nd_parser.h @@ -28,9 +28,9 @@ class TFBatchToSpaceNDParser : public TFNodeParser { TFBatchToSpaceNDParser() = default; ~TFBatchToSpaceNDParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.cc index d160a87fb8..60d15b3200 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFBatchNormParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFBatchNormParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -39,7 +39,7 @@ ops::PrimitiveC *TFBatchNormParser::Parse(const tensorflow::NodeDef &tf_op, for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfBatchNormParser("FusedBatchNormV3", new TFBatchNormParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.h index 15a9b283a3..a40d4fce39 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_batchnorm_parser.h @@ -28,9 +28,9 @@ class TFBatchNormParser : public TFNodeParser { TFBatchNormParser() = default; ~TFBatchNormParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.cc index c41191ae2c..a363fb264e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFBiasAddParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFBiasAddParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFBiasAddParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfBiasAddParser("BiasAdd", new TFBiasAddParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.h index fd5e0557fd..9bd803d91d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_biasadd_parser.h @@ -29,9 +29,9 @@ class TFBiasAddParser : public TFNodeParser { TFBiasAddParser() = default; ~TFBiasAddParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.cc index 57b9b6d8bb..4826844ed6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFBroadcastToParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFBroadcastToParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); if (tf_op.input_size() == 1) { @@ -37,7 +37,7 @@ ops::PrimitiveC *TFBroadcastToParser::Parse(const tensorflow::NodeDef &tf_op, MS_LOG(ERROR) << "add op input failed"; return nullptr; } - return prim.release(); + return prim->GetPrim(); } else { MS_LOG(ERROR) << "broadcast_to has " << tf_op.input_size() << " inputs, invalid"; return nullptr; diff --git a/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.h index 16573ed48b..64da264167 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_broadcast_to_parser.h @@ -28,9 +28,9 @@ class TFBroadcastToParser : public TFNodeParser { TFBroadcastToParser() = default; ~TFBroadcastToParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.cc index e0f27bebe3..c981f5d7d9 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.cc @@ -23,11 +23,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFCastParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFCastParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); auto dst_type = TensorFlowUtils::ParseAttrDataType(tf_op, "DstT"); if (dst_type == kTypeUnknown) { MS_LOG(ERROR) << "Get attr DstT failed"; @@ -36,7 +38,7 @@ ops::PrimitiveC *TFCastParser::Parse(const tensorflow::NodeDef &tf_op, if (dst_type == kNumberTypeInt64) { dst_type = kNumberTypeInt32; } - prim->AddAttr("to", MakeValue(static_cast(dst_type))); + prim_c->AddAttr("to", MakeValue(static_cast(dst_type))); *output_size = 1; if (AddOpInput(tf_op, 0, inputs) != RET_OK) { @@ -44,7 +46,7 @@ ops::PrimitiveC *TFCastParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfCastParser("Cast", new TFCastParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.h index 267ad49abc..abcc9dcdb4 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_cast_parser.h @@ -28,9 +28,9 @@ class TFCastParser : public TFNodeParser { TFCastParser() = default; ~TFCastParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.cc index 0f5065e83a..caa8431092 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.cc @@ -23,11 +23,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFConcatParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFConcatParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); auto axis_node = GetConstInputNode(tf_node_map, tf_op.input(tf_op.input_size() - 1)); if (axis_node == nullptr) { MS_LOG(ERROR) << "get concat axis attr node failed"; @@ -48,8 +50,8 @@ ops::PrimitiveC *TFConcatParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } } - prim->AddAttr(ops::kOriginalOpName, MakeValue("ConcatV2")); - return prim.release(); + prim_c->AddAttr(ops::kOriginalOpName, MakeValue("ConcatV2")); + return prim->GetPrim(); } TFNodeRegistrar g_tfConcatV2Parser("ConcatV2", new TFConcatParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.h index 130ac1ed88..15c2d0b949 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_concat_parser.h @@ -28,9 +28,9 @@ class TFConcatParser : public TFNodeParser { TFConcatParser() = default; ~TFConcatParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.cc index 419be6e22f..84e5bc7bce 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.cc @@ -21,19 +21,22 @@ #include "tools/converter/parser/tf/tf_node_parser_registry.h" #include "tools/converter/parser/tf/tf_util.h" #include "ops/fusion/conv2d_fusion.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { -ops::PrimitiveC *TFConvParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFConvParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); prim->set_pad({0, 0, 0, 0}); prim->set_group(1); auto format = TensorFlowUtils::ParseNodeFormat(tf_op); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); std::vector dilations(2); if (ParseDilations(tf_op, format, &dilations) != RET_OK) { @@ -74,7 +77,7 @@ ops::PrimitiveC *TFConvParser::Parse(const tensorflow::NodeDef &tf_op, } prim->set_pad_list(explicit_paddings); } - prim->AddAttr(ops::kIsOriginalPadMode, MakeValue(is_original_pad_mode)); + prim_c->AddAttr(ops::kIsOriginalPadMode, MakeValue(is_original_pad_mode)); *output_size = 1; if (AddOpInput(tf_op, 0, inputs) != RET_OK || AddOpInput(tf_op, 1, inputs) != RET_OK) { @@ -82,14 +85,14 @@ ops::PrimitiveC *TFConvParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } if (tf_op.op() == "DepthwiseConv2dNative") { - prim->AddAttr(ops::kIsDepthWise, MakeValue(true)); + prim_c->AddAttr(ops::kIsDepthWise, MakeValue(true)); if (prim->GetAttr(ops::kInChannel) != nullptr) { prim->set_group(prim->get_in_channel()); prim->set_out_channel(prim->get_in_channel()); } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfConvParser("Conv2D", new TFConvParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.h index 0a1fe286df..bd9cc6d0bf 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_conv_parser.h @@ -28,9 +28,9 @@ class TFConvParser : public TFConvBaseParser { TFConvParser() = default; ~TFConvParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.cc index ac0b33d2e3..b03638e89a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFCropAndResizeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFCropAndResizeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -56,7 +56,7 @@ ops::PrimitiveC *TFCropAndResizeParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfCropAndResizeParser("CropAndResize", new TFCropAndResizeParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.h index 7901418720..8d65b1531e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_crop_and_resize_parser.h @@ -28,9 +28,9 @@ class TFCropAndResizeParser : public TFNodeParser { TFCropAndResizeParser() = default; ~TFCropAndResizeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.cc index 98d56eb1bf..bef94136de 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.cc @@ -28,15 +28,15 @@ constexpr auto kInputSizeIndex = 0; constexpr auto kFilterIndex = 1; constexpr auto kOutBackpropIndex = 2; -ops::PrimitiveC *TFDeconvParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFDeconvParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_group(1); prim->set_pad({0, 0, 0, 0}); auto format = TensorFlowUtils::ParseNodeFormat(tf_op); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); + prim->AddAttr(mindspore::ops::kOriginalFormat, api::MakeValue(format)); prim->set_output_paddings({0, 0}); std::vector dilations(2); @@ -69,8 +69,8 @@ ops::PrimitiveC *TFDeconvParser::Parse(const tensorflow::NodeDef &tf_op, bool is_original_pad_mode = false; prim->set_pad_mode(ParsePadMode(tf_op, &is_original_pad_mode)); - prim->AddAttr(ops::kIsOriginalPadMode, MakeValue(is_original_pad_mode)); - prim->AddAttr(ops::kOriginalOpName, MakeValue("Conv2DBackpropInput")); + prim->AddAttr(ops::kIsOriginalPadMode, api::MakeValue(is_original_pad_mode)); + prim->AddAttr(ops::kOriginalOpName, api::MakeValue("Conv2DBackpropInput")); *output_size = 1; #ifdef ENABLE_LITE_ACL @@ -85,7 +85,7 @@ ops::PrimitiveC *TFDeconvParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } #endif - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tf_deconv_parser("Conv2DBackpropInput", new TFDeconvParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.h index ccc472e811..67444570d9 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_deconv_parser.h @@ -28,9 +28,9 @@ class TFDeconvParser : public TFConvBaseParser { TFDeconvParser() = default; ~TFDeconvParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.cc index 2cb4f43b60..2ff2c4242a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFDepthToSpaceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFDepthToSpaceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -41,7 +41,7 @@ ops::PrimitiveC *TFDepthToSpaceParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfDepthToSpaceParser("DepthToSpace", new TFDepthToSpaceParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.h index a2c7e63cb9..4c6374aa2a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_depth_to_space_parser.h @@ -28,9 +28,9 @@ class TFDepthToSpaceParser : public TFConvBaseParser { TFDepthToSpaceParser() = default; ~TFDepthToSpaceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.cc index 98154aeda4..09d73fcd3f 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFDropoutParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFDropoutParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -41,7 +41,7 @@ ops::PrimitiveC *TFDropoutParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfDropoutParser("Dropout", new TFDropoutParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.h index ab43592b48..ade96a55ff 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_dropout_parser.h @@ -28,9 +28,9 @@ class TFDropoutParser : public TFNodeParser { TFDropoutParser() = default; ~TFDropoutParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.cc index 786bedacee..eb179b9ae4 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.cc @@ -23,17 +23,17 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFEnterParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { - auto prim = std::make_unique(); +PrimitiveCPtr TFEnterParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfEnterParser("Enter", new TFEnterParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.h index c8c1d25844..2857fa709b 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_enter_parser.h @@ -29,9 +29,9 @@ class TFEnterParser : public TFNodeParser { TFEnterParser() = default; ~TFEnterParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.cc index a2d2cbc550..4692d0761d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.cc @@ -22,17 +22,17 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFExitParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { - auto prim = std::make_unique(); +PrimitiveCPtr TFExitParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfExitParser("Exit", new TFExitParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.h index 9cb2e780df..56d2530433 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_exit_parser.h @@ -29,9 +29,9 @@ class TFExitParser : public TFNodeParser { TFExitParser() = default; ~TFExitParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.cc index 5b557cd26a..98b7bc5327 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFExpandDimsParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFExpandDimsParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFExpandDimsParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfExpandDimsParser("ExpandDims", new TFExpandDimsParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.h index 1f258c150f..ce593fc693 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_expand_dims_parser.h @@ -28,9 +28,9 @@ class TFExpandDimsParser : public TFNodeParser { TFExpandDimsParser() = default; ~TFExpandDimsParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.cc index 539bcc808c..109bc160fb 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFFillParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFFillParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFFillParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfFillParser("Fill", new TFFillParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.h index a6bed38f9e..fa271a176a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_fill_parser.h @@ -29,9 +29,9 @@ class TFFillParser : public TFNodeParser { TFFillParser() = default; ~TFFillParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.cc index b27e3078a6..8e96063dad 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFGatherNDParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFGatherNDParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFGatherNDParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfGatherNDParser("GatherNd", new TFGatherNDParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.h index b288fa4a79..5b53ce2431 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_gather_nd_parser.h @@ -28,9 +28,9 @@ class TFGatherNDParser : public TFNodeParser { TFGatherNDParser() = default; ~TFGatherNDParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.cc index f9e4e7ba8c..457e077294 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.cc @@ -23,11 +23,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFGatherParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFGatherParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); int batchDims = 0; tensorflow::AttrValue attr_value; if (TensorFlowUtils::FindAttrValue(tf_op, "batch_dims", &attr_value)) { @@ -70,7 +72,7 @@ ops::PrimitiveC *TFGatherParser::Parse(const tensorflow::NodeDef &tf_op, if (batchDims != 0 && !axis_is_set) { axis = batchDims; } - prim->AddAttr("axis", MakeValue(axis)); + prim_c->AddAttr("axis", MakeValue(axis)); *output_size = 1; if (AddOpInput(tf_op, 0, inputs) != RET_OK || AddOpInput(tf_op, 1, inputs) != RET_OK) { @@ -78,7 +80,7 @@ ops::PrimitiveC *TFGatherParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfGatherV2Parser("GatherV2", new TFGatherParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.h index 9285c6452a..76e6000a5d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_gather_parser.h @@ -28,9 +28,9 @@ class TFGatherParser : public TFNodeParser { TFGatherParser() = default; ~TFGatherParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_if_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_if_parser.cc index 147535da2c..9ebde67925 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_if_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_if_parser.cc @@ -23,16 +23,16 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFIfParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { - auto prim = std::make_unique(); +PrimitiveCPtr TFIfParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfStatelessIfParser("StatelessIf", new TFIfParser()); TFNodeRegistrar g_tfIfParser("If", new TFIfParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_if_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_if_parser.h index 7be09f4d33..876d71b4a9 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_if_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_if_parser.h @@ -29,9 +29,9 @@ class TFIfParser : public TFNodeParser { TFIfParser() = default; ~TFIfParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.cc index 740a4a66c3..2eaae4fa0d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFInvertPermutationParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFInvertPermutationParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -33,7 +33,7 @@ ops::PrimitiveC *TFInvertPermutationParser::Parse(const tensorflow::NodeDef &tf_ MS_LOG(ERROR) << "Add Op input failed."; return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfInvertPermutationParser("InvertPermutation", new TFInvertPermutationParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.h index e66d5a3f47..8fd175325d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_invert_permutation_parser.h @@ -29,9 +29,9 @@ class TFInvertPermutationParser : public TFNodeParser { TFInvertPermutationParser() = default; ~TFInvertPermutationParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.cc index 992904a2c9..eca5896507 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFIsFiniteParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFIsFiniteParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -33,7 +33,7 @@ ops::PrimitiveC *TFIsFiniteParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tf_is_finite_parser("IsFinite", new TFIsFiniteParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.h index 911c0181be..ec0839e927 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_is_finite_parser.h @@ -29,9 +29,9 @@ class TFIsFiniteParser : public TFNodeParser { TFIsFiniteParser() = default; ~TFIsFiniteParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.cc index 99c97cd2e3..ddef051038 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFLinSpaceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLinSpaceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF LinSpaceParser"; auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -37,7 +37,7 @@ ops::PrimitiveC *TFLinSpaceParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfLinSpaceParser("LinSpace", new TFLinSpaceParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.h index 6dc2cb1c83..e956074637 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_linspace_parser.h @@ -28,9 +28,9 @@ class TFLinSpaceParser : public TFNodeParser { TFLinSpaceParser() = default; ~TFLinSpaceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.cc index 42c651ee2c..e18c207439 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.cc @@ -25,9 +25,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFLogicalAndParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLogicalAndParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -35,12 +35,12 @@ ops::PrimitiveC *TFLogicalAndParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFLogicalOrParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLogicalOrParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -48,12 +48,12 @@ ops::PrimitiveC *TFLogicalOrParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFLogicalNotParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLogicalNotParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -61,7 +61,7 @@ ops::PrimitiveC *TFLogicalNotParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfLogicalAndParser("LogicalAnd", new TFLogicalAndParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.h index 95437ced53..029c747ab6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_logical_parser.h @@ -29,9 +29,9 @@ class TFLogicalAndParser : public TFNodeParser { TFLogicalAndParser() = default; ~TFLogicalAndParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFLogicalOrParser : public TFNodeParser { @@ -39,9 +39,9 @@ class TFLogicalOrParser : public TFNodeParser { TFLogicalOrParser() = default; ~TFLogicalOrParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFLogicalNotParser : public TFNodeParser { @@ -49,9 +49,9 @@ class TFLogicalNotParser : public TFNodeParser { TFLogicalNotParser() = default; ~TFLogicalNotParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.cc index 53c8f1122f..1226ce584e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFLoopCondParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFLoopCondParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); @@ -32,7 +32,7 @@ ops::PrimitiveC *TFLoopCondParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfLoopCondParser("LoopCond", new TFLoopCondParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.h index fd4cc2358a..ef0177b448 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_loop_cond_parser.h @@ -29,9 +29,9 @@ class TFLoopCondParser : public TFNodeParser { TFLoopCondParser() = default; ~TFLoopCondParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.cc index 05b6f8b18e..d4e4f2f59a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFMatMulParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFMatMulParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -43,7 +43,7 @@ ops::PrimitiveC *TFMatMulParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfMatMulParser("MatMul", new TFMatMulParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.h index 986332f6af..635d7f9cdc 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_matmul_parser.h @@ -29,9 +29,9 @@ class TFMatMulParser : public TFNodeParser { TFMatMulParser() = default; ~TFMatMulParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.cc index 828c5be97f..26584817d6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.cc @@ -23,17 +23,17 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFMergeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { - auto prim = std::make_unique(); +PrimitiveCPtr TFMergeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfMergeParser("Merge", new TFMergeParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.h index dccd6f7260..87d32ad493 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_merge_parser.h @@ -29,9 +29,9 @@ class TFMergeParser : public TFNodeParser { TFMergeParser() = default; ~TFMergeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc index a907d2bc88..271d5330c4 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_model_parser.cc @@ -319,6 +319,11 @@ STATUS SetStringTensorInfo(const tensorflow::TensorProto &tensor_proto, tensor:: delete tensor_data; return RET_OK; } + +FuncGraphPtr ConvertGraph(api::FuncGraphPtr func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} } // namespace STATUS TFModelParser::ConvertConstVariant(const tensorflow::TensorProto &tensor_proto, tensor::TensorPtr *tensor_info) { @@ -536,14 +541,16 @@ api::FuncGraphPtr TFModelParser::Parse(const converter::ConverterParameters &fla ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; } - res_graph_ = std::make_shared(); + auto graph = std::make_shared(); + MS_CHECK_TRUE_MSG(graph != nullptr, nullptr, "create FuncGraph failed"); + res_graph_ = api::MakeShared(graph); if (res_graph_ == nullptr) { MS_LOG(ERROR) << "funGraphPtr is nullptr"; ReturnCode::GetSingleReturnCode()->UpdateReturnCode(RET_ERROR); return nullptr; } - res_graph_->set_attr("graph_name", MakeValue("main_graph")); - res_graph_->set_attr("fmk", MakeValue(static_cast(converter::kFmkTypeTf))); + graph->set_attr("graph_name", MakeValue("main_graph")); + graph->set_attr("fmk", MakeValue(static_cast(converter::kFmkTypeTf))); for (int i = 0; i < tf_root_graph_->node_size(); i++) { auto &node_def = tf_root_graph_->node(i); @@ -551,9 +558,7 @@ api::FuncGraphPtr TFModelParser::Parse(const converter::ConverterParameters &fla tf_root_graph_nodes_vec_.emplace_back(&node_def); } - auto func_graph = std::dynamic_pointer_cast(res_graph_); - MS_CHECK_TRUE_RET(func_graph != nullptr, nullptr); - status = ConvertGraphInputsAndConsts(tf_root_graph_nodes_vec_, func_graph, &anf_root_node_map_, true); + status = ConvertGraphInputsAndConsts(tf_root_graph_nodes_vec_, graph, &anf_root_node_map_, true); if (status != RET_OK) { ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; @@ -561,7 +566,7 @@ api::FuncGraphPtr TFModelParser::Parse(const converter::ConverterParameters &fla bool success_flag = true; for (int i = 0; i < tf_root_graph_->node_size(); i++) { auto &node_def = tf_root_graph_->node(i); - status = ConvertOps(node_def, tf_root_graph_nodes_, func_graph, &anf_root_node_map_); + status = ConvertOps(node_def, tf_root_graph_nodes_, graph, &anf_root_node_map_); ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); if (status != RET_OK) { success_flag = false; @@ -595,13 +600,13 @@ api::FuncGraphPtr TFModelParser::Parse(const converter::ConverterParameters &fla return nullptr; } - if ((status = CommonAnfAdjust(func_graph)) != RET_OK) { + if ((status = CommonAnfAdjust(graph)) != RET_OK) { MS_LOG(ERROR) << "AdjustForAnf failed."; ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); return nullptr; } std::set all_func_graphs = {}; - GetAllFuncGraph(func_graph, &all_func_graphs); + GetAllFuncGraph(graph, &all_func_graphs); if ((status = TF2AnfAdjust(all_func_graphs)) != RET_OK) { MS_LOG(ERROR) << "TF2AnfAdjust failed."; ReturnCode::GetSingleReturnCode()->UpdateReturnCode(status); @@ -609,12 +614,12 @@ api::FuncGraphPtr TFModelParser::Parse(const converter::ConverterParameters &fla } auto unify_format = std::make_shared(kFmkTypeTf, false); MS_CHECK_TRUE_RET(unify_format != nullptr, nullptr); - if (!unify_format->Run(func_graph)) { + if (!unify_format->Run(graph)) { MS_LOG(ERROR) << "Run insert transpose failed."; return nullptr; } - func_graph->set_manager(nullptr); - static auto root_func_manager = Manage(func_graph); + graph->set_manager(nullptr); + static auto root_func_manager = Manage(graph); return res_graph_; } @@ -806,7 +811,7 @@ STATUS TFModelParser::ControlFlowNodePostProcess(const std::map(res_graph_); + auto func_graph = ConvertGraph(res_graph_); if (func_graph == nullptr) { MS_LOG(ERROR) << "func graph is invalid."; return RET_ERROR; @@ -828,7 +833,7 @@ STATUS TFModelParser::ControlFlowNodePostProcess(const std::mapinputs(); inputs.insert(inputs.begin() + 1, {first_value_node, second_value_node}); - auto new_node = res_graph_->NewCNode(inputs); // must create new node, otherwise node_users won't update + auto new_node = func_graph->NewCNode(inputs); // must create new node, otherwise node_users won't update if (new_node == nullptr) { MS_LOG(ERROR) << "new node failed"; return RET_ERROR; @@ -907,7 +912,9 @@ STATUS TFModelParser::ConvertOutputTensor(const tensorflow::NodeDef &op, const C MS_LOG(ERROR) << "new TupleGetItem failed"; return RET_NULL_PTR; } - auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_ptr); + auto prim_c = tuple_get_item_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(prim_c); + auto tuple_get_item_prim = NewValueNode(prim_c); CHECK_NULL_RETURN(tuple_get_item_prim); auto get_item_value = NewValueNode(MakeValue(output_idx)); CHECK_NULL_RETURN(get_item_value); @@ -965,12 +972,12 @@ STATUS TFModelParser::ConvertOps(const tensorflow::NodeDef &node_def, return RET_OK; } MS_LOG(INFO) << "parse op : " << op_type; - ops::PrimitiveC *primitive_c; + ops::PrimitiveCPtr primitive_c; auto node_parser = registry::NodeParserRegistry::GetNodeParser(kFmkTypeTf, op_type); int output_size; std::vector input_names; if (node_parser != nullptr) { - primitive_c = node_parser->Parse(node_def, tf_node_map, &input_names, &output_size); + primitive_c = node_parser->Parse(node_def, tf_node_map, &input_names, &output_size)->GetPrim(); } else { auto node_parser_builtin = TFNodeParserRegistry::GetInstance()->GetNodeParser(op_type); if (node_parser_builtin == nullptr) { @@ -989,7 +996,7 @@ STATUS TFModelParser::ConvertOps(const tensorflow::NodeDef &node_def, for (int i = 0; i < output_size; i++) { node_output_num_[node_def.name() + ":" + to_string(i)] = 1; } - auto value_node = NewValueNode(std::shared_ptr(primitive_c)); + auto value_node = NewValueNode(primitive_c); if (value_node == nullptr) { MS_LOG(ERROR) << "value_node is nullptr"; return RET_ERROR; @@ -1067,7 +1074,7 @@ STATUS TFModelParser::ProcessControlFlowOp(const CNodePtr &anf_node, const strin } STATUS TFModelParser::ConvertQuantParams(const size_t &input_size, const size_t &output_size, - ops::PrimitiveC *primitive_c) { + PrimitiveCPtr primitive_c) { if (primitive_c == nullptr) { MS_LOG(ERROR) << "primitive_c is null, get quant params failed."; return RET_NULL_PTR; @@ -1142,7 +1149,7 @@ STATUS TFModelParser::ConvertRootGraphOutputs() { MS_LOG(ERROR) << "get graph outputs node error"; return status; } - auto func_graph = std::dynamic_pointer_cast(res_graph_); + auto func_graph = ConvertGraph(res_graph_); if (func_graph == nullptr) { MS_LOG(ERROR) << "unc graph is invalid."; return RET_ERROR; @@ -1168,7 +1175,9 @@ STATUS TFModelParser::MakeAnfGraphOutputs(const std::vector &output_ MS_LOG(ERROR) << "new MakeTuple failed"; return RET_NULL_PTR; } - auto make_tuple_prim = NewValueNode(make_tuple_prim_ptr); + auto make_tuple_prim_c = make_tuple_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(make_tuple_prim_c); + auto make_tuple_prim = NewValueNode(make_tuple_prim_c); CHECK_NULL_RETURN(make_tuple_prim); make_tuple_inputs.insert(make_tuple_inputs.begin(), make_tuple_prim); auto make_tuple_cnode = anf_graph->NewCNode(make_tuple_inputs); @@ -1180,7 +1189,9 @@ STATUS TFModelParser::MakeAnfGraphOutputs(const std::vector &output_ MS_LOG(ERROR) << "new Return failed"; return RET_NULL_PTR; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(return_prim_c); + auto value_node = NewValueNode(return_prim_c); CHECK_NULL_RETURN(value_node); std::vector op_inputs = {value_node, make_tuple_cnode}; auto cnode = anf_graph->NewCNode(op_inputs); @@ -1193,7 +1204,9 @@ STATUS TFModelParser::MakeAnfGraphOutputs(const std::vector &output_ MS_LOG(ERROR) << "new Return failed"; return RET_NULL_PTR; } - auto value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + CHECK_NULL_RETURN(return_prim_c); + auto value_node = NewValueNode(return_prim_c); CHECK_NULL_RETURN(value_node); std::vector op_inputs{value_node, output_nodes.front()}; auto return_cnode = anf_graph->NewCNode(op_inputs); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_model_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_model_parser.h index 5f264f28be..f8dbf60451 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_model_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_model_parser.h @@ -92,7 +92,7 @@ class TFModelParser : public converter::ModelParser { STATUS ControlFlowNodePostProcess(const std::map &first_func_map, const std::map &second_func_map); - static STATUS ConvertQuantParams(const size_t &input_size, const size_t &output_size, ops::PrimitiveC *primitive_c); + static STATUS ConvertQuantParams(const size_t &input_size, const size_t &output_size, PrimitiveCPtr primitive_c); static STATUS MakeAnfGraphOutputs(const std::vector &output_nodes, const FuncGraphPtr &anf_graph); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.cc index 7bfa222e59..b348378073 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.cc @@ -23,17 +23,16 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFNegParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFNegParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfNegParser("Neg", new TFNegParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.h index 2bc0ab32b1..bd49fd9677 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_neg_parser.h @@ -29,9 +29,9 @@ class TFNegParser : public TFNodeParser { TFNegParser() = default; ~TFNegParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.cc index 2673326966..e3e70abdc2 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.cc @@ -22,17 +22,17 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFNextIterationParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { - auto prim = std::make_unique(); +PrimitiveCPtr TFNextIterationParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfNextIterationParser("NextIteration", new TFNextIterationParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.h index 2cfa9d601d..cda133451b 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_next_iteration_parser.h @@ -29,9 +29,9 @@ class TFNextIterationParser : public TFNodeParser { TFNextIterationParser() = default; ~TFNextIterationParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_node_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_node_parser.h index dc0c39463a..f9e6a5ce14 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_node_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_node_parser.h @@ -26,6 +26,9 @@ #include "ops/primitive_c.h" #include "mindspore/core/utils/check_convert_utils.h" #include "nnacl/op_base.h" +#include "tools/converter/parser/parser_utils.h" +#include "ops/op_utils.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { @@ -35,9 +38,9 @@ class TFNodeParser { virtual ~TFNodeParser() = default; - virtual ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { + virtual PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { return nullptr; } diff --git a/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.cc index e60f4dca2f..77f839c298 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFNonMaxSuppressionParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFNonMaxSuppressionParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_center_point_box(0); @@ -38,7 +38,7 @@ ops::PrimitiveC *TFNonMaxSuppressionParser::Parse(const tensorflow::NodeDef &tf_ } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfNonMaxSuppressionV3Parser("NonMaxSuppressionV3", new TFNonMaxSuppressionParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.h index 4915ffe6da..5af064f69c 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_non_max_suppression_parser.h @@ -29,9 +29,9 @@ class TFNonMaxSuppressionParser : public TFNodeParser { TFNonMaxSuppressionParser() = default; ~TFNonMaxSuppressionParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.cc index 8bb6d00d95..91258084e3 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFOneHotParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFOneHotParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -42,7 +42,7 @@ ops::PrimitiveC *TFOneHotParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfOneHotParser("OneHot", new TFOneHotParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.h index 0ceab33391..435bbb76b6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_one_hot_parser.h @@ -28,9 +28,9 @@ class TFOneHotParser : public TFNodeParser { TFOneHotParser() = default; ~TFOneHotParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.cc index 5058312dd3..fccfc4cdf6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFPackParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFPackParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -43,7 +43,7 @@ ops::PrimitiveC *TFPackParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfPackParser("Pack", new TFPackParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.h index d5630f6ef7..8d54184ad1 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_pack_parser.h @@ -28,9 +28,9 @@ class TFPackParser : public TFNodeParser { TFPackParser() = default; ~TFPackParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.cc index 54d1e526cb..03017f6f56 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.cc @@ -28,15 +28,17 @@ namespace { constexpr int kInputIndexTwo = 2; constexpr int kInputSizeThree = 3; } // namespace -ops::PrimitiveC *TFPadParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFPadParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); if (tf_op.op() == "Pad") { prim->set_padding_mode(mindspore::PaddingMode::CONSTANT); prim->set_constant_value(0.0f); - prim->AddAttr(ops::kOriginalOpName, MakeValue("Pad")); + prim_c->AddAttr(ops::kOriginalOpName, MakeValue("Pad")); } else if (tf_op.op() == "PadV2") { prim->set_padding_mode(mindspore::PaddingMode::CONSTANT); if (tf_op.input_size() < kInputSizeThree) { @@ -59,7 +61,7 @@ ops::PrimitiveC *TFPadParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } prim->set_constant_value(tensor_proto.float_val(0)); - prim->AddAttr(ops::kOriginalOpName, MakeValue("PadV2")); + prim_c->AddAttr(ops::kOriginalOpName, MakeValue("PadV2")); } else if (tf_op.op() == "MirrorPad") { tensorflow::AttrValue attr_value; if (!TensorFlowUtils::FindAttrValue(tf_op, "mode", &attr_value)) { @@ -75,7 +77,7 @@ ops::PrimitiveC *TFPadParser::Parse(const tensorflow::NodeDef &tf_op, MS_LOG(ERROR) << "padding mode:" << attr_value.s() << " don't support"; return nullptr; } - prim->AddAttr(ops::kOriginalOpName, MakeValue("MirrorPad")); + prim_c->AddAttr(ops::kOriginalOpName, MakeValue("MirrorPad")); } *output_size = 1; @@ -84,7 +86,7 @@ ops::PrimitiveC *TFPadParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfPadParser("Pad", new TFPadParser()); TFNodeRegistrar g_tfPadV2Parser("PadV2", new TFPadParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.h index e3aaa95567..b786ea1d71 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_pad_parser.h @@ -28,9 +28,9 @@ class TFPadParser : public TFNodeParser { TFPadParser() = default; ~TFPadParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.cc index 971bf44e0e..776ae22e62 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.cc @@ -27,11 +27,13 @@ namespace mindspore { namespace lite { constexpr int kTfPoolStrideListSize = 4; constexpr int kTfPoolKernelListSize = 4; -ops::PrimitiveC *TFMaxPoolParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFMaxPoolParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); tensorflow::AttrValue attr_value; if (TensorFlowUtils::FindAttrValue(tf_op, "padding", &attr_value)) { if (attr_value.s() == "VALID") { @@ -42,7 +44,7 @@ ops::PrimitiveC *TFMaxPoolParser::Parse(const tensorflow::NodeDef &tf_op, } auto format = TensorFlowUtils::ParseNodeFormat(tf_op); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); if (TensorFlowUtils::FindAttrValue(tf_op, "strides", &attr_value)) { const auto &stride_list = attr_value.list(); @@ -69,14 +71,15 @@ ops::PrimitiveC *TFMaxPoolParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TFAvgPoolParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFAvgPoolParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); - + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); tensorflow::AttrValue attr_value; if (TensorFlowUtils::FindAttrValue(tf_op, "padding", &attr_value)) { if (attr_value.s() == "VALID") { @@ -87,7 +90,7 @@ ops::PrimitiveC *TFAvgPoolParser::Parse(const tensorflow::NodeDef &tf_op, } auto format = TensorFlowUtils::ParseNodeFormat(tf_op); - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(format)); if (TensorFlowUtils::FindAttrValue(tf_op, "strides", &attr_value)) { const auto &stride_list = attr_value.list(); @@ -114,7 +117,7 @@ ops::PrimitiveC *TFAvgPoolParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfMaxPoolParser("MaxPool", new TFMaxPoolParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.h index 646d7f1f76..0b0b70a7c5 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_pool_parser.h @@ -28,9 +28,9 @@ class TFMaxPoolParser : public TFNodeParser { TFMaxPoolParser() = default; ~TFMaxPoolParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; class TFAvgPoolParser : public TFNodeParser { @@ -38,9 +38,9 @@ class TFAvgPoolParser : public TFNodeParser { TFAvgPoolParser() = default; ~TFAvgPoolParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.cc index 19dc4bf236..729bcfe4e6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFRaggedRangeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFRaggedRangeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { *output_size = 2; for (int i = 0; i < 3; i++) { if (AddOpInput(tf_op, i, inputs) != RET_OK) { @@ -34,7 +34,7 @@ ops::PrimitiveC *TFRaggedRangeParser::Parse(const tensorflow::NodeDef &tf_op, } auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfRaggedRangeParser("RaggedRange", new TFRaggedRangeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.h index efdd3699a6..2aa3adc285 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_ragged_range_parser.h @@ -27,9 +27,9 @@ class TFRaggedRangeParser : public TFNodeParser { TFRaggedRangeParser() = default; ~TFRaggedRangeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.cc index ebe24cfd16..443ff51d32 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFRandomStandardNormalParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFRandomStandardNormalParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -46,7 +46,7 @@ ops::PrimitiveC *TFRandomStandardNormalParser::Parse(const tensorflow::NodeDef & return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfRandomStandardNormalParser("RandomStandardNormal", new TFRandomStandardNormalParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.h index c1fd43499d..9dd56bcd4f 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_random_standard_normal_parser.h @@ -28,9 +28,9 @@ class TFRandomStandardNormalParser : public TFNodeParser { TFRandomStandardNormalParser() = default; ~TFRandomStandardNormalParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_range_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_range_parser.cc index b049ebe1a7..cf2c99a504 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_range_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_range_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFRangeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFRangeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -85,7 +85,7 @@ ops::PrimitiveC *TFRangeParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfRangeParser("Range", new TFRangeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_range_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_range_parser.h index decd7cbbf6..b08fdac13e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_range_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_range_parser.h @@ -28,9 +28,9 @@ class TFRangeParser : public TFNodeParser { TFRangeParser() = default; ~TFRangeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.cc index 3c98975a9d..2b4aacef2e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFRankParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFRankParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF RankParser"; if (output_size == nullptr) { MS_LOG(ERROR) << "output_size is nullptr"; @@ -41,7 +41,7 @@ ops::PrimitiveC *TFRankParser::Parse(const tensorflow::NodeDef &tf_op, if (status != RET_OK) { return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfRankParser("Rank", new TFRankParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.h index efea02bd8d..6c75a2831b 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_rank_parser.h @@ -28,9 +28,9 @@ class TFRankParser : public TFNodeParser { TFRankParser() = default; ~TFRankParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.cc index 72e199538a..4b8d9afea2 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFReduceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFReduceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); if (tf_op.op() == "Sum") { @@ -63,7 +63,7 @@ ops::PrimitiveC *TFReduceParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSumParser("Sum", new TFReduceParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.h index 3c3411654b..e6cc581535 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_reduce_parser.h @@ -28,9 +28,9 @@ class TFReduceParser : public TFNodeParser { TFReduceParser() = default; ~TFReduceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.cc index 2bf9409f6a..af0c3bcd84 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFReshapeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFReshapeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFReshapeParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfReshapeParser("Reshape", new TFReshapeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.h index 4d99a77c8e..f6b78c057a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_reshape_parser.h @@ -28,9 +28,9 @@ class TFReshapeParser : public TFNodeParser { TFReshapeParser() = default; ~TFReshapeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.cc index f0a23206c1..7a5262e954 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.cc @@ -24,13 +24,15 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFResizeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFResizeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); tensorflow::AttrValue attr_value; - prim->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NHWC)); + prim_c->AddAttr(mindspore::ops::kOriginalFormat, MakeValue(mindspore::Format::NHWC)); prim->set_cubic_coeff(-0.75f); if (!TensorFlowUtils::FindAttrValue(tf_op, "align_corners", &attr_value)) { MS_LOG(ERROR) << "The align_corners attr should be specified"; @@ -38,11 +40,11 @@ ops::PrimitiveC *TFResizeParser::Parse(const tensorflow::NodeDef &tf_op, } if (attr_value.b()) { prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::ALIGN_CORNERS); - prim->AddAttr("align_corners", MakeValue(true)); + prim_c->AddAttr("align_corners", MakeValue(true)); } else if (TensorFlowUtils::FindAttrValue(tf_op, "half_pixel_centers", &attr_value) && attr_value.b()) { prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::HALF_PIXEL); prim->set_cubic_coeff(-0.5f); - prim->AddAttr("half_pixel_centers", MakeValue(true)); + prim_c->AddAttr("half_pixel_centers", MakeValue(true)); } else { prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::ASYMMETRIC); } @@ -77,7 +79,7 @@ ops::PrimitiveC *TFResizeParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfResizeBilinearParser("ResizeBilinear", new TFResizeParser()); TFNodeRegistrar g_tfResizeNearestNeighborParser("ResizeNearestNeighbor", new TFResizeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.h index 359e00c34e..0917402819 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_resize_parser.h @@ -28,9 +28,9 @@ class TFResizeParser : public TFNodeParser { TFResizeParser() = default; ~TFResizeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.cc index 6b3928342b..083631eac6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFReverseParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFReverseParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -60,7 +60,7 @@ ops::PrimitiveC *TFReverseParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfReverseV2Parser("ReverseV2", new TFReverseParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.h index 43dafd9b67..ec0e21c6bb 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_reverse_parser.h @@ -29,9 +29,9 @@ class TFReverseParser : public TFNodeParser { TFReverseParser() = default; ~TFReverseParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.cc index edeccf7d02..de5ae735ae 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFReverseSequenceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFReverseSequenceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -46,7 +46,7 @@ ops::PrimitiveC *TFReverseSequenceParser::Parse(const tensorflow::NodeDef &tf_op return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfReverseSequenceParser("ReverseSequence", new TFReverseSequenceParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.h index 229e83b551..18b637c607 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_reverse_sequence_parser.h @@ -28,9 +28,9 @@ class TFReverseSequenceParser : public TFNodeParser { TFReverseSequenceParser() = default; ~TFReverseSequenceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_select_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_select_parser.cc index d8eb081e97..9374a564ad 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_select_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_select_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSelectParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSelectParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -33,7 +33,7 @@ ops::PrimitiveC *TFSelectParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSelectParser("Select", new TFSelectParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_select_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_select_parser.h index a771378eae..c97141d4e2 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_select_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_select_parser.h @@ -29,9 +29,9 @@ class TFSelectParser : public TFNodeParser { TFSelectParser() = default; ~TFSelectParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.cc index 1dc91828d5..f927973904 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFShapeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFShapeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFShapeParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfShapeParser("Shape", new TFShapeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.h index d0b0799e7c..3b90ede137 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_shape_parser.h @@ -28,9 +28,9 @@ class TFShapeParser : public TFNodeParser { TFShapeParser() = default; ~TFShapeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_size_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_size_parser.cc index 7584af78cb..c905c067d3 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_size_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_size_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSizeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSizeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFSizeParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSizeParser("Size", new TFSizeParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_size_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_size_parser.h index c31c1025a0..24ff511bfb 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_size_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_size_parser.h @@ -29,9 +29,9 @@ class TFSizeParser : public TFNodeParser { TFSizeParser() = default; ~TFSizeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.cc index e6c2512ace..a172628b2c 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSliceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSliceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); // begin @@ -68,7 +68,7 @@ ops::PrimitiveC *TFSliceParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSliceParser("Slice", new TFSliceParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.h index e2d170dae3..01903d72ed 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_slice_parser.h @@ -29,9 +29,9 @@ class TFSliceParser : public TFNodeParser { TFSliceParser() = default; ~TFSliceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.cc index 6015f14d2c..08ed1ec466 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSoftmaxParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSoftmaxParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -41,7 +41,7 @@ ops::PrimitiveC *TFSoftmaxParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSoftmaxParser("Softmax", new TFSoftmaxParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.h index 8c46590eca..6cb3789d2e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_softmax_parser.h @@ -28,9 +28,9 @@ class TFSoftmaxParser : public TFNodeParser { TFSoftmaxParser() = default; ~TFSoftmaxParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.cc index bb815f7fc3..b2ac193b2e 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSpaceToBatchNDParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSpaceToBatchNDParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -35,7 +35,7 @@ ops::PrimitiveC *TFSpaceToBatchNDParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSpaceToBatchNDParser("SpaceToBatchND", new TFSpaceToBatchNDParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.h index d00bb00ea9..214c4bf06d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_space_to_batch_nd_parser.h @@ -28,9 +28,9 @@ class TFSpaceToBatchNDParser : public TFNodeParser { TFSpaceToBatchNDParser() = default; ~TFSpaceToBatchNDParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_split_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_split_parser.cc index 4c6417b0f7..cb8e6aa7d5 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_split_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_split_parser.cc @@ -25,9 +25,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSplitParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSplitParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -93,7 +93,7 @@ ops::PrimitiveC *TFSplitParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSplitParser("Split", new TFSplitParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_split_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_split_parser.h index 9f33008021..c75ca4b33d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_split_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_split_parser.h @@ -28,9 +28,9 @@ class TFSplitParser : public TFNodeParser { TFSplitParser() = default; ~TFSplitParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.cc index cccebd365c..7cf98f338f 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSqueezeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSqueezeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); std::vector axis; @@ -46,7 +46,7 @@ ops::PrimitiveC *TFSqueezeParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSqueezeParser("Squeeze", new TFSqueezeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.h index f5998dcfc5..c79c0920fa 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_squeeze_parser.h @@ -28,9 +28,9 @@ class TFSqueezeParser : public TFNodeParser { TFSqueezeParser() = default; ~TFSqueezeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.cc index c84d67cd42..94bb11f4ac 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFStrideSliceParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFStrideSliceParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -67,7 +67,7 @@ ops::PrimitiveC *TFStrideSliceParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfStrideSliceParser("StridedSlice", new TFStrideSliceParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.h index 03fdaad661..9943d05b49 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_stride_slice_parser.h @@ -28,9 +28,9 @@ class TFStrideSliceParser : public TFNodeParser { TFStrideSliceParser() = default; ~TFStrideSliceParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.cc index e58424e347..0e2faaa145 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFSwitchParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFSwitchParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 2; @@ -33,7 +33,7 @@ ops::PrimitiveC *TFSwitchParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfSwitchParser("Switch", new TFSwitchParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.h index 9f15e157cb..f379659a30 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_switch_parser.h @@ -29,9 +29,9 @@ class TFSwitchParser : public TFNodeParser { TFSwitchParser() = default; ~TFSwitchParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.cc index e034294cfa..24de968240 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.cc @@ -23,11 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorArrayGatherParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorArrayGatherParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF TensorArrayGatherParser"; - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr"; return nullptr; @@ -36,7 +36,7 @@ ops::PrimitiveC *TFTensorArrayGatherParser::Parse(const tensorflow::NodeDef &tf_ for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfTensorArrayGatherParser("TensorArrayGatherV3", new TFTensorArrayGatherParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.h index 76ab611611..f039a71dec 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_gather_parser.h @@ -28,9 +28,9 @@ class TFTensorArrayGatherParser : public TFNodeParser { TFTensorArrayGatherParser() = default; ~TFTensorArrayGatherParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.cc index 290bc7208b..034cf3ccee 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.cc @@ -23,11 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorArrayParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorArrayParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF TensorArrayParser"; - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr"; return nullptr; @@ -38,7 +38,7 @@ ops::PrimitiveC *TFTensorArrayParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfTensorArrayParser("TensorArrayV3", new TFTensorArrayParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.h index ddc303c73a..973d826c47 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_parser.h @@ -28,9 +28,9 @@ class TFTensorArrayParser : public TFNodeParser { TFTensorArrayParser() = default; ~TFTensorArrayParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.cc index b865eb2ab3..d6cac1a5ae 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.cc @@ -23,11 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorArrayReadParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorArrayReadParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF TensorArrayReadParser"; - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr"; return nullptr; @@ -37,7 +37,7 @@ ops::PrimitiveC *TFTensorArrayReadParser::Parse(const tensorflow::NodeDef &tf_op for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfTensorArrayReadParser("TensorArrayReadV3", new TFTensorArrayReadParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.h index 287b8ae744..57ef217255 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_read_parser.h @@ -27,9 +27,9 @@ class TFTensorArrayReadParser : public TFNodeParser { public: TFTensorArrayReadParser() = default; ~TFTensorArrayReadParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.cc index 7e7287e150..4b02fb3f3b 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.cc @@ -23,11 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorArrayScatterParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorArrayScatterParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF TensorArrayScatterParser"; - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr"; return nullptr; @@ -37,7 +37,7 @@ ops::PrimitiveC *TFTensorArrayScatterParser::Parse(const tensorflow::NodeDef &tf for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfTensorArrayScatterParser("TensorArrayScatterV3", new TFTensorArrayScatterParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.h index ebbf796029..e0ec0bca24 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_scatter_parser.h @@ -27,9 +27,9 @@ class TFTensorArrayScatterParser : public TFNodeParser { public: TFTensorArrayScatterParser() = default; ~TFTensorArrayScatterParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.cc index 1c8af81c0d..0f1bd83f22 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.cc @@ -23,11 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorArraySizeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorArraySizeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF TensorArraySizeParser"; - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr"; return nullptr; @@ -37,7 +37,7 @@ ops::PrimitiveC *TFTensorArraySizeParser::Parse(const tensorflow::NodeDef &tf_op inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfTensorArraySizeParser("TensorArraySizeV3", new TFTensorArraySizeParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.h index 02e967775d..ac7eeb1962 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_size_parser.h @@ -27,9 +27,9 @@ class TFTensorArraySizeParser : public TFNodeParser { public: TFTensorArraySizeParser() = default; ~TFTensorArraySizeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.cc index e4091dc0c1..c4343a5e12 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.cc @@ -23,11 +23,11 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorArrayWriteParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorArrayWriteParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF TensorArrayWriteParser"; - auto prim = std::make_unique(); + auto prim = std::make_shared(); if (prim == nullptr) { MS_LOG(ERROR) << "prim is nullptr"; return nullptr; @@ -38,7 +38,7 @@ ops::PrimitiveC *TFTensorArrayWriteParser::Parse(const tensorflow::NodeDef &tf_o inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfTensorArrayWriteParser("TensorArrayWriteV3", new TFTensorArrayWriteParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.h index 39dafec2f6..4d99191957 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_array_write_parser.h @@ -28,9 +28,9 @@ class TFTensorArrayWriteParser : public TFNodeParser { TFTensorArrayWriteParser() = default; ~TFTensorArrayWriteParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.cc index d398cfaa67..773684b8dd 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorListFromTensorParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorListFromTensorParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -59,7 +59,7 @@ ops::PrimitiveC *TFTensorListFromTensorParser::Parse(const tensorflow::NodeDef & } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTensorListFromTensorParser("TensorListFromTensor", new TFTensorListFromTensorParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.h index 49f950367f..6e745354a2 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_from_tensor_parser.h @@ -28,9 +28,9 @@ class TFTensorListFromTensorParser : public TFNodeParser { TFTensorListFromTensorParser() = default; ~TFTensorListFromTensorParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.cc index 8e78b8f788..2e09ac2d57 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorListGetItemParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorListGetItemParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -48,7 +48,7 @@ ops::PrimitiveC *TFTensorListGetItemParser::Parse(const tensorflow::NodeDef &tf_ } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTensorListGetItemParser("TensorListGetItem", new TFTensorListGetItemParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.h index f3b8224b93..cb089bec5d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_get_item_parser.h @@ -29,9 +29,9 @@ class TFTensorListGetItemParser : public TFNodeParser { TFTensorListGetItemParser() = default; ~TFTensorListGetItemParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.cc index 80028215e1..a8a646b7c8 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorListReserveParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorListReserveParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -59,7 +59,7 @@ ops::PrimitiveC *TFTensorListReserveParser::Parse(const tensorflow::NodeDef &tf_ } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTensorListReserveParser("TensorListReserve", new TFTensorListReserveParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.h index 4b2ce85433..e4e255f9c3 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_reserve_parser.h @@ -28,9 +28,9 @@ class TFTensorListReserveParser : public TFNodeParser { TFTensorListReserveParser() = default; ~TFTensorListReserveParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.cc index a756c335b8..8107df8c3a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorListSetItemParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorListSetItemParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -47,7 +47,7 @@ ops::PrimitiveC *TFTensorListSetItemParser::Parse(const tensorflow::NodeDef &tf_ return nullptr; } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTensorListSetItemParser("TensorListSetItem", new TFTensorListSetItemParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.h index 5e7dde35c6..a7c6a20ece 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_set_item_parser.h @@ -28,9 +28,9 @@ class TFTensorListSetItemParser : public TFNodeParser { TFTensorListSetItemParser() = default; ~TFTensorListSetItemParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.cc index fa029dbba2..b2957eb043 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTensorListStackParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTensorListStackParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -54,7 +54,7 @@ ops::PrimitiveC *TFTensorListStackParser::Parse(const tensorflow::NodeDef &tf_op } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTensorListStackParser("TensorListStack", new TFTensorListStackParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.h index c47a4bc1c2..4eafa3fbd9 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tensor_list_stack_parser.h @@ -28,9 +28,9 @@ class TFTensorListStackParser : public TFNodeParser { TFTensorListStackParser() = default; ~TFTensorListStackParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.cc index 94dfc70873..279f47c3b6 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTileParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTileParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -48,7 +48,7 @@ ops::PrimitiveC *TFTileParser::Parse(const tensorflow::NodeDef &tf_op, MS_LOG(ERROR) << "add op input failed"; return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTileParser("Tile", new TFTileParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.h index 587face512..75893cc415 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_tile_parser.h @@ -28,9 +28,9 @@ class TFTileParser : public TFNodeParser { TFTileParser() = default; ~TFTileParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.cc index 75287e9a83..095d25fb3a 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTopKParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTopKParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -40,8 +40,7 @@ ops::PrimitiveC *TFTopKParser::Parse(const tensorflow::NodeDef &tf_op, MS_LOG(ERROR) << "Add Op input failed."; return nullptr; } - - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTopKV2Parser("TopKV2", new TFTopKParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.h index 3fafb4aebd..70ce985eaf 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_topk_parser.h @@ -29,9 +29,9 @@ class TFTopKParser : public TFNodeParser { TFTopKParser() = default; ~TFTopKParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.cc index 5114700d42..4d0a767f7f 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFTransposeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFTransposeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = 1; @@ -34,7 +34,7 @@ ops::PrimitiveC *TFTransposeParser::Parse(const tensorflow::NodeDef &tf_op, return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfTransposeParser("Transpose", new TFTransposeParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.h index cfa6d78fb2..1ac8b9172d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_transpose_parser.h @@ -28,9 +28,9 @@ class TFTransposeParser : public TFNodeParser { TFTransposeParser() = default; ~TFTransposeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.cc index fc5fbeb4e1..07236eb61b 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFUniformRealParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFUniformRealParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { MS_LOG(DEBUG) << "TF UniformRealParser"; auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -46,7 +46,7 @@ ops::PrimitiveC *TFUniformRealParser::Parse(const tensorflow::NodeDef &tf_op, MS_LOG(ERROR) << "add op input failed"; return nullptr; } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfRandomUniformParser("RandomUniform", new TFUniformRealParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.h index 77177a5389..6fbf083443 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_uniform_real_parser.h @@ -28,9 +28,9 @@ class TFUniformRealParser : public TFNodeParser { TFUniformRealParser() = default; ~TFUniformRealParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.cc index 5660066a0a..dded3109ab 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFUnpackParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFUnpackParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); tensorflow::AttrValue attr_value; @@ -43,7 +43,7 @@ ops::PrimitiveC *TFUnpackParser::Parse(const tensorflow::NodeDef &tf_op, } } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfUnpackParser("Unpack", new TFUnpackParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.h index 930d2b7567..2ac70131c8 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_unpack_parser.h @@ -28,9 +28,9 @@ class TFUnpackParser : public TFNodeParser { TFUnpackParser() = default; ~TFUnpackParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_where_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_where_parser.cc index 63c3dc4b53..d733ca1a04 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_where_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_where_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFWhereParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFWhereParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); @@ -33,7 +33,7 @@ ops::PrimitiveC *TFWhereParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfWhereParser("Where", new TFWhereParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_where_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_where_parser.h index e669af5066..7b77d74e7b 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_where_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_where_parser.h @@ -29,9 +29,9 @@ class TFWhereParser : public TFNodeParser { TFWhereParser() = default; ~TFWhereParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_while_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_while_parser.cc index 7cdb1de110..cbcdbbf857 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_while_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_while_parser.cc @@ -23,17 +23,17 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFWhileParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { - auto prim = std::make_unique(); +PrimitiveCPtr TFWhileParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); for (int i = 0; i < tf_op.input_size(); i++) { inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim; } TFNodeRegistrar g_tfStatelessWhileParser("StatelessWhile", new TFWhileParser()); diff --git a/mindspore/lite/tools/converter/parser/tf/tf_while_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_while_parser.h index 7de0c1880d..de3a59ef41 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_while_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_while_parser.h @@ -29,9 +29,9 @@ class TFWhileParser : public TFNodeParser { TFWhileParser() = default; ~TFWhileParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.cc b/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.cc index 49e94006a6..75bc6edd0d 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.cc +++ b/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TFZerosLikeParser::Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) { +PrimitiveCPtr TFZerosLikeParser::Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); *output_size = tf_op.input_size(); @@ -33,7 +33,7 @@ ops::PrimitiveC *TFZerosLikeParser::Parse(const tensorflow::NodeDef &tf_op, inputs->emplace_back(tf_op.input(i)); } - return prim.release(); + return prim->GetPrim(); } TFNodeRegistrar g_tfZerosLikeParser("ZerosLike", new TFZerosLikeParser()); } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.h b/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.h index b09c85e2b7..86145de354 100644 --- a/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.h +++ b/mindspore/lite/tools/converter/parser/tf/tf_zeros_like_parser.h @@ -28,9 +28,9 @@ class TFZerosLikeParser : public TFNodeParser { TFZerosLikeParser() = default; ~TFZerosLikeParser() override = default; - ops::PrimitiveC *Parse(const tensorflow::NodeDef &tf_op, - const std::map &tf_node_map, - std::vector *inputs, int *output_size) override; + PrimitiveCPtr Parse(const tensorflow::NodeDef &tf_op, + const std::map &tf_node_map, + std::vector *inputs, int *output_size) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tf_bidirection_gru_cf_fusion.cc b/mindspore/lite/tools/converter/parser/tf_bidirection_gru_cf_fusion.cc index b3c0bbf859..de0d5eec7d 100644 --- a/mindspore/lite/tools/converter/parser/tf_bidirection_gru_cf_fusion.cc +++ b/mindspore/lite/tools/converter/parser/tf_bidirection_gru_cf_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/parser/tf_bidirection_gru_cf_fusion.h" #include #include diff --git a/mindspore/lite/tools/converter/parser/tflite/schema.fbs b/mindspore/lite/tools/converter/parser/tflite/schema.fbs index a8bdf5e067..1919b8bea1 100644 --- a/mindspore/lite/tools/converter/parser/tflite/schema.fbs +++ b/mindspore/lite/tools/converter/parser/tflite/schema.fbs @@ -1,1094 +1,1094 @@ -// Copyright 2017 The TensorFlow Authors. All Rights Reserved. -// -// 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. - -// Revision History -// Version 0: Initial version. -// Version 1: Add subgraphs to schema. -// Version 2: Rename operators to conform to NN API. -// Version 3: Move buffer data from Model.Subgraph.Tensors to Model.Buffers. - -namespace tflite; - -// This corresponds to the version. -file_identifier "TFL3"; -// File extension of any written files. -file_extension "tflite"; - -// IMPORTANT: All new members of tables, enums and unions must be added at the -// end to ensure backwards compatibility. - -// The type of data stored in a tensor. -enum TensorType : byte { - FLOAT32 = 0, - FLOAT16 = 1, - INT32 = 2, - UINT8 = 3, - INT64 = 4, - STRING = 5, - BOOL = 6, - INT16 = 7, - COMPLEX64 = 8, - INT8 = 9, - FLOAT64 = 10, -} - -// Custom quantization parameters for experimenting with new quantization -// techniques. -table CustomQuantization { - custom:[ubyte] (force_align: 16); -} - -// Represents a specific quantization technique's parameters. -union QuantizationDetails { - CustomQuantization, -} - -// Parameters for converting a quantized tensor back to float. -table QuantizationParameters { - // These four parameters are the asymmetric linear quantization parameters. - // Given a quantized value q, the corresponding float value f should be: - // f = scale * (q - zero_point) - // For other quantization types, the QuantizationDetails below is used. - min:[float]; // For importing back into tensorflow. - max:[float]; // For importing back into tensorflow. - scale:[float]; // For dequantizing the tensor's values. - zero_point:[long]; - - // If this is not none, the other quantization parameters (i.e. min, max, - // scale, zero_point fields above) are ignored and the value of the - // QuantizationDetails union should be used. - details:QuantizationDetails; - - // Specifies the dimension of the Tensor's shape that the scales and - // zero_points correspond to. For example, a tensor t, with dims=[4, 3, 2, 1] - // with quantization params: - // scale=[1.0, 2.0, 3.0], zero_point=[1, 2, 3], quantization_dimension=1 - // will be quantized across the second dimension of t. - // t[:, 0, :, :] will have scale[0]=1.0, zero_point[0]=1 - // t[:, 1, :, :] will have scale[1]=2.0, zero_point[0]=2 - // t[:, 2, :, :] will have scale[2]=3.0, zero_point[0]=3 - quantized_dimension:int; -} - -// Sparse tensors. -// We use a modification of the TACO format. -// Reference: http://tensor-compiler.org/kjolstad-oopsla17-tensor-compiler.pdf -// -// To encode a conceptual n-dimensional dense tensor with dims (d0, ..., dn-1), -// potentially with a k-dimensional block (0 <= k <= n) with dims -// (dn, ..., dn+k-1), the format needs to specify: -// 1. In what order to traverse these dimensions. For example, to store a 2-D -// matrix in row major order, the traversal order would be (d0, d1), -// whereas to store it in column major order, the traversal order would be -// (d1, d0). If the 2-D matrix has a 2-D inner block, the traversal order -// could be (d0, d1, d2, d3). -// 2. How each block dimension in (dn, ..., dn+k-1) maps to the original -// tensor dimension in (d0, ..., dn-1). -// 3. In the traversal order defined above, the format (dense vs. sparse) and -// index metadata for each dimension. For a dense dimension, this is just -// the size of that dimension. For a sparse dimension, it's the same as -// the compressed index defined in the Compressed Sparse Row (CSR) format. -// (http://scipy-lectures.org/advanced/scipy_sparse/csr_matrix.html) - -// The storage type for a dimension. Currently we support: -// 1. DENSE: each coordinate in this dimension is stored implicitly. -// 2. SPARSE_CSR: only the coordinates with non-zero elements are stored. The -// compression technique is the same what CSR uses. -// More types like a sparse dimension with a different compression technique -// could be added to the list in the future. -enum DimensionType : byte { - DENSE = 0, - SPARSE_CSR = 1, -} - -table Int32Vector { - values:[int]; -} - -table Uint16Vector { - values:[ushort] (force_align: 4); -} - -table Uint8Vector { - values:[ubyte] (force_align: 4); -} - -// Variable-typed buffer to store the index metadata for a sparse dimension. -// The widest type is Int32 instead of UInt32 because tensor's shape is a int32 -// vector. We don't want the per-dimensional index to overflow that range. -union SparseIndexVector { - Int32Vector, - Uint16Vector, - Uint8Vector -} - -table DimensionMetadata { - // Whether a dimension is dense or sparse. - format:DimensionType; - // Index metadata used for a dimension. - // - If format is DimensionType.DENSE then we use the dense_size field to - // store the size of that dimension. Each index in that dimension is - // stored implicitly. - // - If format is DimensionType.SPARSE_CSR then we use array_segments and - // array_indices to encode that dimension. array_segments represents how - // to segment the indices array, each segment corresponds to one element - // in the previous dimension. array_indices represents the index of the - // non-zero elements within this dimension (as those in the CSR matrix - // format, where the first array is row pointers and the second array is - // column indices). - dense_size:int; - array_segments:SparseIndexVector; - array_indices:SparseIndexVector; -} - -// Parameters to encode a sparse TfLite tensor. -table SparsityParameters { - // The traversal order of the dimensions defined in the `shape` field of the - // conceptual dense tensor. For a n-dimensional tensors with dims (d0, d1, - // ..., dn-1), - // - if not block sparse, the traversal_order is just a permutation of (d0, - // ..., dn-1). For example, a 2-D matrix stored in row-major order would - // have traversal_order = (d0, d1). - // - if block sparse with a k-dimensional block (0 <= k <= n), the - // traversal_order has n + k elements. The first n elements are still a - // permutation of (d0, ..., dn-1). The lask k elements are a permutation - // of (dn, ..., dn+k-1), defining how to traverse a block internally. For - // example, a 2-D matrix with 2-D blocks, both stored in row-major order - // would have traversal_order = (d0, d1, d2, d3). - traversal_order:[int]; - // For an n-dimensional tensor with a k-dimensional block (0 <= k <= n), - // stores how a block dimension in (dn, ..., dn+k-1) maps to the original - // tensor dimension in (d0, ..., dn). - // It's stored in the order of (dn, ..., dn+k-1). - // If not block-sparse, this field is NULL. - block_map:[int]; - // In the traversal order defined above, the metadata needed for - // each dimension to locate the non-zero values in the original dense tensor. - // The size of the dim_metadata array = the size of the traversal_order array - // = n + k. - dim_metadata:[DimensionMetadata]; -} - -table Tensor { - // The tensor shape. The meaning of each entry is operator-specific but - // builtin ops use: [batch size, height, width, number of channels] (That's - // Tensorflow's NHWC). - shape:[int]; - type:TensorType; - // An index that refers to the buffers table at the root of the model. Or, - // if there is no data buffer associated (i.e. intermediate results), then - // this is 0 (which refers to an always existent empty buffer). - // - // The data_buffer itself is an opaque container, with the assumption that the - // target device is little-endian. In addition, all builtin operators assume - // the memory is ordered such that if `shape` is [4, 3, 2], then index - // [i, j, k] maps to data_buffer[i*3*2 + j*2 + k]. - buffer:uint; - name:string; // For debugging and importing back into tensorflow. - quantization:QuantizationParameters; // Optional. - - is_variable:bool = false; - - // Parameters to encode a sparse tensor. See the example in - // tensorflow/lite/testdata/sparse_tensor.json. - sparsity:SparsityParameters; // Optional. - - // Encodes `shape` with unknown dimensions. Unknown dimensions are - // represented with -1. - shape_signature:[int]; // Optional. -} - -// A list of builtin operators. Builtin operators are slightly faster than custom -// ones, but not by much. Moreover, while custom operators accept an opaque -// object containing configuration parameters, builtins have a predetermined -// set of acceptable options. - -enum BuiltinOperator : byte { - ADD = 0, - AVERAGE_POOL_2D = 1, - CONCATENATION = 2, - CONV_2D = 3, - DEPTHWISE_CONV_2D = 4, - DEPTH_TO_SPACE = 5, - DEQUANTIZE = 6, - EMBEDDING_LOOKUP = 7, - FLOOR = 8, - FULLY_CONNECTED = 9, - HASHTABLE_LOOKUP = 10, - L2_NORMALIZATION = 11, - L2_POOL_2D = 12, - LOCAL_RESPONSE_NORMALIZATION = 13, - LOGISTIC = 14, - LSH_PROJECTION = 15, - LSTM = 16, - MAX_POOL_2D = 17, - MUL = 18, - RELU = 19, - // NOTE(aselle): RELU_N1_TO_1 used to be called RELU1, but it was renamed - // since different model developers use RELU1 in different ways. Never - // create another op called RELU1. - RELU_N1_TO_1 = 20, - RELU6 = 21, - RESHAPE = 22, - RESIZE_BILINEAR = 23, - RNN = 24, - SOFTMAX = 25, - SPACE_TO_DEPTH = 26, - SVDF = 27, - TANH = 28, - // Consider rename to CONCATENATE_EMBEDDINGS - CONCAT_EMBEDDINGS = 29, - SKIP_GRAM = 30, - CALL = 31, - CUSTOM = 32, - EMBEDDING_LOOKUP_SPARSE = 33, - PAD = 34, - UNIDIRECTIONAL_SEQUENCE_RNN = 35, - GATHER = 36, - BATCH_TO_SPACE_ND = 37, - SPACE_TO_BATCH_ND = 38, - TRANSPOSE = 39, - MEAN = 40, - SUB = 41, - DIV = 42, - SQUEEZE = 43, - UNIDIRECTIONAL_SEQUENCE_LSTM = 44, - STRIDED_SLICE = 45, - BIDIRECTIONAL_SEQUENCE_RNN = 46, - EXP = 47, - TOPK_V2 = 48, - SPLIT = 49, - LOG_SOFTMAX = 50, - // DELEGATE is a special op type for the operations which are delegated to - // other backends. - // WARNING: Experimental interface, subject to change - DELEGATE = 51, - BIDIRECTIONAL_SEQUENCE_LSTM = 52, - CAST = 53, - PRELU = 54, - MAXIMUM = 55, - ARG_MAX = 56, - MINIMUM = 57, - LESS = 58, - NEG = 59, - PADV2 = 60, - GREATER = 61, - GREATER_EQUAL = 62, - LESS_EQUAL = 63, - SELECT = 64, - SLICE = 65, - SIN = 66, - TRANSPOSE_CONV = 67, - SPARSE_TO_DENSE = 68, - TILE = 69, - EXPAND_DIMS = 70, - EQUAL = 71, - NOT_EQUAL = 72, - LOG = 73, - SUM = 74, - SQRT = 75, - RSQRT = 76, - SHAPE = 77, - POW = 78, - ARG_MIN = 79, - FAKE_QUANT = 80, - REDUCE_PROD = 81, - REDUCE_MAX = 82, - PACK = 83, - LOGICAL_OR = 84, - ONE_HOT = 85, - LOGICAL_AND = 86, - LOGICAL_NOT = 87, - UNPACK = 88, - REDUCE_MIN = 89, - FLOOR_DIV = 90, - REDUCE_ANY = 91, - SQUARE = 92, - ZEROS_LIKE = 93, - FILL = 94, - FLOOR_MOD = 95, - RANGE = 96, - RESIZE_NEAREST_NEIGHBOR = 97, - LEAKY_RELU = 98, - SQUARED_DIFFERENCE = 99, - MIRROR_PAD = 100, - ABS = 101, - SPLIT_V = 102, - UNIQUE = 103, - CEIL = 104, - REVERSE_V2 = 105, - ADD_N = 106, - GATHER_ND = 107, - COS = 108, - WHERE = 109, - RANK = 110, - ELU = 111, - REVERSE_SEQUENCE = 112, - MATRIX_DIAG = 113, - QUANTIZE = 114, - MATRIX_SET_DIAG = 115, - ROUND = 116, - HARD_SWISH = 117, - IF = 118, - WHILE = 119, - NON_MAX_SUPPRESSION_V4 = 120, - NON_MAX_SUPPRESSION_V5 = 121, - SCATTER_ND = 122, - SELECT_V2 = 123, - DENSIFY = 124, - SEGMENT_SUM = 125, - BATCH_MATMUL = 126 -} - - -// Options for the builtin operators. -union BuiltinOptions { - Conv2DOptions, - DepthwiseConv2DOptions, - ConcatEmbeddingsOptions, - LSHProjectionOptions, - Pool2DOptions, - SVDFOptions, - RNNOptions, - FullyConnectedOptions, - SoftmaxOptions, - ConcatenationOptions, - AddOptions, - L2NormOptions, - LocalResponseNormalizationOptions, - LSTMOptions, - ResizeBilinearOptions, - CallOptions, - ReshapeOptions, - SkipGramOptions, - SpaceToDepthOptions, - EmbeddingLookupSparseOptions, - MulOptions, - PadOptions, - GatherOptions, - BatchToSpaceNDOptions, - SpaceToBatchNDOptions, - TransposeOptions, - ReducerOptions, - SubOptions, - DivOptions, - SqueezeOptions, - SequenceRNNOptions, - StridedSliceOptions, - ExpOptions, - TopKV2Options, - SplitOptions, - LogSoftmaxOptions, - CastOptions, - DequantizeOptions, - MaximumMinimumOptions, - ArgMaxOptions, - LessOptions, - NegOptions, - PadV2Options, - GreaterOptions, - GreaterEqualOptions, - LessEqualOptions, - SelectOptions, - SliceOptions, - TransposeConvOptions, - SparseToDenseOptions, - TileOptions, - ExpandDimsOptions, - EqualOptions, - NotEqualOptions, - ShapeOptions, - PowOptions, - ArgMinOptions, - FakeQuantOptions, - PackOptions, - LogicalOrOptions, - OneHotOptions, - LogicalAndOptions, - LogicalNotOptions, - UnpackOptions, - FloorDivOptions, - SquareOptions, - ZerosLikeOptions, - FillOptions, - BidirectionalSequenceLSTMOptions, - BidirectionalSequenceRNNOptions, - UnidirectionalSequenceLSTMOptions, - FloorModOptions, - RangeOptions, - ResizeNearestNeighborOptions, - LeakyReluOptions, - SquaredDifferenceOptions, - MirrorPadOptions, - AbsOptions, - SplitVOptions, - UniqueOptions, - ReverseV2Options, - AddNOptions, - GatherNdOptions, - CosOptions, - WhereOptions, - RankOptions, - ReverseSequenceOptions, - MatrixDiagOptions, - QuantizeOptions, - MatrixSetDiagOptions, - HardSwishOptions, - IfOptions, - WhileOptions, - DepthToSpaceOptions, - NonMaxSuppressionV4Options, - NonMaxSuppressionV5Options, - ScatterNdOptions, - SelectV2Options, - DensifyOptions, - SegmentSumOptions, - BatchMatMulOptions -} - -enum Padding : byte { SAME, VALID } - -enum ActivationFunctionType : byte { - NONE = 0, - RELU = 1, - RELU_N1_TO_1 = 2, - RELU6 = 3, - TANH = 4, - SIGN_BIT = 5, -} - -table Conv2DOptions { - padding:Padding; - stride_w:int; - stride_h:int; - fused_activation_function:ActivationFunctionType; - dilation_w_factor:int = 1; - dilation_h_factor:int = 1; -} - -table Pool2DOptions { - padding:Padding; - stride_w:int; - stride_h:int; - filter_width:int; - filter_height:int; - fused_activation_function:ActivationFunctionType; -} - -table DepthwiseConv2DOptions { - // Parameters for DepthwiseConv version 1 or above. - padding:Padding; - stride_w:int; - stride_h:int; - // `depth_multiplier` is redundant. It's used by CPU kernels in - // TensorFlow 2.0 or below, but ignored in versions above. - // See comments in lite/c/builtin_op_data.h for more details. - depth_multiplier:int; - fused_activation_function:ActivationFunctionType; - // Parameters for DepthwiseConv version 2 or above. - dilation_w_factor:int = 1; - dilation_h_factor:int = 1; -} - -table ConcatEmbeddingsOptions { - num_channels:int; - num_columns_per_channel:[int]; - embedding_dim_per_channel:[int]; // This could be inferred from parameters. -} - -enum LSHProjectionType: byte { - UNKNOWN = 0, - SPARSE = 1, - DENSE = 2, -} - -table LSHProjectionOptions { - type: LSHProjectionType; -} - -table SVDFOptions { - rank:int; - fused_activation_function:ActivationFunctionType; - // For weights-only quantization, use asymmetric quantization for non - // constant inputs at evaluation time. - asymmetric_quantize_inputs:bool; -} - -// An implementation of TensorFlow RNNCell. -table RNNOptions { - fused_activation_function:ActivationFunctionType; - asymmetric_quantize_inputs:bool; -} - -// An implementation of TensorFlow dynamic_rnn with RNNCell. -table SequenceRNNOptions { - time_major:bool; - fused_activation_function:ActivationFunctionType; - asymmetric_quantize_inputs:bool; -} - -// An implementation of TensorFlow bidrectional_dynamic_rnn with RNNCell. -table BidirectionalSequenceRNNOptions { - time_major:bool; - fused_activation_function:ActivationFunctionType; - merge_outputs: bool; - asymmetric_quantize_inputs:bool; -} - -enum FullyConnectedOptionsWeightsFormat: byte { - DEFAULT = 0, - SHUFFLED4x16INT8 = 1, -} - -// An implementation of TensorFlow fully_connected (a.k.a Dense) layer. -table FullyConnectedOptions { - // Parameters for FullyConnected version 1 or above. - fused_activation_function:ActivationFunctionType; - - // Parameters for FullyConnected version 2 or above. - weights_format:FullyConnectedOptionsWeightsFormat = DEFAULT; - - // Parameters for FullyConnected version 5 or above. - // If set to true, then the number of dimension is preserved. Furthermore, - // all but the last dimension of the input and output shapes will be equal. - keep_num_dims: bool; - - // Parameters for FullyConnected version 7 or above. - // If set to true, then weights-only op will use asymmetric quantization for - // inputs. - asymmetric_quantize_inputs: bool; -} - -table SoftmaxOptions { - beta: float; -} - -// An implementation of TensorFlow concat. -table ConcatenationOptions { - axis:int; - fused_activation_function:ActivationFunctionType; -} - -table AddOptions { - fused_activation_function:ActivationFunctionType; -} - -table MulOptions { - fused_activation_function:ActivationFunctionType; -} - -table L2NormOptions { - fused_activation_function:ActivationFunctionType; -} - -table LocalResponseNormalizationOptions { - radius:int; - bias:float; - alpha:float; - beta:float; -} - -enum LSTMKernelType : byte { - // Full LSTM kernel which supports peephole and projection. - FULL = 0, - // Basic LSTM kernels. Equivalent to TensorFlow BasicLSTMCell. - BASIC = 1, -} - -// An implementation of TensorFlow LSTMCell and CoupledInputForgetGateLSTMCell -table LSTMOptions { - // Parameters for LSTM version 1 or above. - fused_activation_function:ActivationFunctionType; - cell_clip: float; // Optional, 0.0 means no clipping - proj_clip: float; // Optional, 0.0 means no clipping - - // Parameters for LSTM version 2 or above. - // Basic kernel is only supported in version 2 or above. - kernel_type: LSTMKernelType = FULL; - - // Parameters for LSTM version 4 or above. - asymmetric_quantize_inputs: bool; -} - -// An implementation of TensorFlow dynamic_rnn with LSTMCell. -table UnidirectionalSequenceLSTMOptions { - fused_activation_function:ActivationFunctionType; - cell_clip: float; // Optional, 0.0 means no clipping - proj_clip: float; // Optional, 0.0 means no clipping - - // If true then first dimension is sequence, otherwise batch. - time_major:bool; - - // Parameter for Unidirectional Sequence LSTM version 4. - asymmetric_quantize_inputs:bool; -} - -table BidirectionalSequenceLSTMOptions { - // Parameters supported by version 1: - fused_activation_function:ActivationFunctionType; - cell_clip: float; // Optional, 0.0 means no clipping - proj_clip: float; // Optional, 0.0 means no clipping - - // If true, store the outputs of both directions into the first output. - merge_outputs: bool; - - // Parameters supported by version 2: - // If true then first dimension is sequence, otherwise batch. - // Version 1 implementations assumed time_major to be true, so this default - // value should never change. - time_major: bool = true; - - // Parameters for version 3 or above. - asymmetric_quantize_inputs:bool; -} - -table ResizeBilinearOptions { - new_height: int (deprecated); - new_width: int (deprecated); - align_corners: bool; - half_pixel_centers: bool; -} - -table ResizeNearestNeighborOptions { - align_corners: bool; - half_pixel_centers: bool; -} - -// A call operation options -table CallOptions { - // The subgraph index that needs to be called. - subgraph:uint; -} - -table PadOptions { -} - -table PadV2Options { -} - -table ReshapeOptions { - new_shape:[int]; -} - -table SpaceToBatchNDOptions { -} - -table BatchToSpaceNDOptions { -} - -table SkipGramOptions { - ngram_size: int; - max_skip_size: int; - include_all_ngrams: bool; -} - -table SpaceToDepthOptions { - block_size: int; -} - -table DepthToSpaceOptions { - block_size: int; -} - -table SubOptions { - fused_activation_function:ActivationFunctionType; -} - -table DivOptions { - fused_activation_function:ActivationFunctionType; -} - -table TopKV2Options { -} - -enum CombinerType : byte { - SUM = 0, - MEAN = 1, - SQRTN = 2, -} - -table EmbeddingLookupSparseOptions { - combiner:CombinerType; -} - -table GatherOptions { - axis: int; -} - -table TransposeOptions { -} - -table ExpOptions { -} - -table CosOptions { -} - -table ReducerOptions { - keep_dims: bool; -} - -table SqueezeOptions { - squeeze_dims:[int]; -} - -table SplitOptions { - num_splits: int; -} - -table SplitVOptions { - num_splits: int; -} - -table StridedSliceOptions { - begin_mask: int; - end_mask: int; - ellipsis_mask: int; - new_axis_mask: int; - shrink_axis_mask: int; -} - -table LogSoftmaxOptions { -} - -table CastOptions { - in_data_type: TensorType; - out_data_type: TensorType; -} - -table DequantizeOptions { -} - -table MaximumMinimumOptions { -} - -table TileOptions { -} - -table ArgMaxOptions { - output_type : TensorType; -} - -table ArgMinOptions { - output_type : TensorType; -} - -table GreaterOptions { -} - -table GreaterEqualOptions { -} - -table LessOptions { -} - -table LessEqualOptions { -} - -table NegOptions { -} - -table SelectOptions { -} - -table SliceOptions { -} - -table TransposeConvOptions { - padding:Padding; - stride_w:int; - stride_h:int; -} - -table ExpandDimsOptions { -} - -table SparseToDenseOptions { - validate_indices:bool; -} - -table EqualOptions { -} - -table NotEqualOptions { -} - -table ShapeOptions { - // Optional output type of the operation (int32 or int64). Defaults to int32. - out_type : TensorType; -} - -table RankOptions { -} - -table PowOptions { -} - -table FakeQuantOptions { - // Parameters supported by version 1: - min:float; - max:float; - num_bits:int; - - // Parameters supported by version 2: - narrow_range:bool; -} - -table PackOptions { - values_count:int; - axis:int; -} - -table LogicalOrOptions { -} - -table OneHotOptions { - axis:int; -} - -table AbsOptions { -} - - -table HardSwishOptions { -} - -table LogicalAndOptions { -} - -table LogicalNotOptions { -} - -table UnpackOptions { - num:int; - axis:int; -} - -table FloorDivOptions { -} - -table SquareOptions { -} - -table ZerosLikeOptions { -} - -table FillOptions { -} - -table FloorModOptions { -} - -table RangeOptions { -} - -table LeakyReluOptions { - alpha:float; -} - -table SquaredDifferenceOptions { -} - -enum MirrorPadMode : byte { - // Doesn't include borders. - REFLECT = 0, - // Includes borders. - SYMMETRIC = 1, -} - -table MirrorPadOptions { - mode:MirrorPadMode; -} - -table UniqueOptions { - idx_out_type:TensorType = INT32; -} - -table ReverseV2Options { -} - -table AddNOptions { -} - -table GatherNdOptions { -} - -table WhereOptions { -} - -table ReverseSequenceOptions { - seq_dim:int; - batch_dim:int = 0; -} - -table MatrixDiagOptions { -} - -table QuantizeOptions { -} - -table MatrixSetDiagOptions { -} - -table IfOptions { - then_subgraph_index:int; - else_subgraph_index:int; -} - -table WhileOptions { - cond_subgraph_index:int; - body_subgraph_index:int; -} - -table NonMaxSuppressionV4Options { -} - -table NonMaxSuppressionV5Options { -} - -table ScatterNdOptions { -} - -table SelectV2Options { -} - -table DensifyOptions { -} - -table SegmentSumOptions { -} - -table BatchMatMulOptions { - adj_x:bool; - adj_y:bool; -} - -// An OperatorCode can be an enum value (BuiltinOperator) if the operator is a -// builtin, or a string if the operator is custom. -table OperatorCode { - builtin_code:BuiltinOperator; - custom_code:string; - - // The version of the operator. The version need to be bumped whenever new - // parameters are introduced into an op. - version:int = 1; -} - -enum CustomOptionsFormat : byte { - FLEXBUFFERS = 0, -} - -// An operator takes tensors as inputs and outputs. The type of operation being -// performed is determined by an index into the list of valid OperatorCodes, -// while the specifics of each operations is configured using builtin_options -// or custom_options. -table Operator { - // Index into the operator_codes array. Using an integer here avoids - // complicate map lookups. - opcode_index:uint; - - // Optional input are indicated by -1. - inputs:[int]; - outputs:[int]; - - builtin_options:BuiltinOptions; - custom_options:[ubyte]; - custom_options_format:CustomOptionsFormat; - - // A list of booleans indicating the input tensors which are being mutated by - // this operator.(e.g. used by RNN and LSTM). - // For example, if the "inputs" array refers to 5 tensors and the second and - // fifth are mutable variables, then this list will contain - // [false, true, false, false, true]. - // - // If the list is empty, no variable is mutated in this operator. - // The list either has the same length as `inputs`, or is empty. - mutating_variable_inputs:[bool]; - - // A list of indices to the subgraph's "tensors" that are internal to an Op. - // Internal tensors are those that do not flow in or out of the operation, - // but instead are part of internal computation. As such, the operation's - // implementation may manage its memory more efficiently. They are needed - // however (i.e. not just an implementation detail) since they are part of the - // computation, which may require relevant metadata such as quantization - // parameters. - intermediates:[int]; -} - -// The root type, defining a subgraph, which typically represents an entire -// model. -table SubGraph { - // A list of all tensors used in this subgraph. - tensors:[Tensor]; - - // Indices of the tensors that are inputs into this subgraph. Note this is - // the list of non-static tensors that feed into the subgraph for inference. - inputs:[int]; - - // Indices of the tensors that are outputs out of this subgraph. Note this is - // the list of output tensors that are considered the product of the - // subgraph's inference. - outputs:[int]; - - // All operators, in execution order. - operators:[Operator]; - - // Name of this subgraph (used for debugging). - name:string; -} - -// Table of raw data buffers (used for constant tensors). Referenced by tensors -// by index. The generous alignment accommodates mmap-friendly data structures. -table Buffer { - data:[ubyte] (force_align: 16); -} - -table Metadata { - // A human readable string to uniquely identify a Metadata. - name:string; - // An index to the buffers table. - buffer:uint; -} - -table Model { - // Version of the schema. - version:uint; - - // A list of all operator codes used in this model. This is - // kept in order because operators carry an index into this - // vector. - operator_codes:[OperatorCode]; - - // All the subgraphs of the model. The 0th is assumed to be the main - // model. - subgraphs:[SubGraph]; - - // A description of the model. - description:string; - - // Buffers of the model. - // Note the 0th entry of this array must be an empty buffer (sentinel). - // This is a convention so that tensors without a buffer can provide 0 as - // their buffer. - buffers:[Buffer]; - - // Metadata about the model. Indirects into the existings buffers list. - // Deprecated, prefer to use metadata field. - metadata_buffer:[int]; - - // Metadata about the model. - metadata:[Metadata]; -} - -root_type Model; +// Copyright 2017 The TensorFlow Authors. All Rights Reserved. +// +// 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. + +// Revision History +// Version 0: Initial version. +// Version 1: Add subgraphs to schema. +// Version 2: Rename operators to conform to NN API. +// Version 3: Move buffer data from Model.Subgraph.Tensors to Model.Buffers. + +namespace tflite; + +// This corresponds to the version. +file_identifier "TFL3"; +// File extension of any written files. +file_extension "tflite"; + +// IMPORTANT: All new members of tables, enums and unions must be added at the +// end to ensure backwards compatibility. + +// The type of data stored in a tensor. +enum TensorType : byte { + FLOAT32 = 0, + FLOAT16 = 1, + INT32 = 2, + UINT8 = 3, + INT64 = 4, + STRING = 5, + BOOL = 6, + INT16 = 7, + COMPLEX64 = 8, + INT8 = 9, + FLOAT64 = 10, +} + +// Custom quantization parameters for experimenting with new quantization +// techniques. +table CustomQuantization { + custom:[ubyte] (force_align: 16); +} + +// Represents a specific quantization technique's parameters. +union QuantizationDetails { + CustomQuantization, +} + +// Parameters for converting a quantized tensor back to float. +table QuantizationParameters { + // These four parameters are the asymmetric linear quantization parameters. + // Given a quantized value q, the corresponding float value f should be: + // f = scale * (q - zero_point) + // For other quantization types, the QuantizationDetails below is used. + min:[float]; // For importing back into tensorflow. + max:[float]; // For importing back into tensorflow. + scale:[float]; // For dequantizing the tensor's values. + zero_point:[long]; + + // If this is not none, the other quantization parameters (i.e. min, max, + // scale, zero_point fields above) are ignored and the value of the + // QuantizationDetails union should be used. + details:QuantizationDetails; + + // Specifies the dimension of the Tensor's shape that the scales and + // zero_points correspond to. For example, a tensor t, with dims=[4, 3, 2, 1] + // with quantization params: + // scale=[1.0, 2.0, 3.0], zero_point=[1, 2, 3], quantization_dimension=1 + // will be quantized across the second dimension of t. + // t[:, 0, :, :] will have scale[0]=1.0, zero_point[0]=1 + // t[:, 1, :, :] will have scale[1]=2.0, zero_point[0]=2 + // t[:, 2, :, :] will have scale[2]=3.0, zero_point[0]=3 + quantized_dimension:int; +} + +// Sparse tensors. +// We use a modification of the TACO format. +// Reference: http://tensor-compiler.org/kjolstad-oopsla17-tensor-compiler.pdf +// +// To encode a conceptual n-dimensional dense tensor with dims (d0, ..., dn-1), +// potentially with a k-dimensional block (0 <= k <= n) with dims +// (dn, ..., dn+k-1), the format needs to specify: +// 1. In what order to traverse these dimensions. For example, to store a 2-D +// matrix in row major order, the traversal order would be (d0, d1), +// whereas to store it in column major order, the traversal order would be +// (d1, d0). If the 2-D matrix has a 2-D inner block, the traversal order +// could be (d0, d1, d2, d3). +// 2. How each block dimension in (dn, ..., dn+k-1) maps to the original +// tensor dimension in (d0, ..., dn-1). +// 3. In the traversal order defined above, the format (dense vs. sparse) and +// index metadata for each dimension. For a dense dimension, this is just +// the size of that dimension. For a sparse dimension, it's the same as +// the compressed index defined in the Compressed Sparse Row (CSR) format. +// (http://scipy-lectures.org/advanced/scipy_sparse/csr_matrix.html) + +// The storage type for a dimension. Currently we support: +// 1. DENSE: each coordinate in this dimension is stored implicitly. +// 2. SPARSE_CSR: only the coordinates with non-zero elements are stored. The +// compression technique is the same what CSR uses. +// More types like a sparse dimension with a different compression technique +// could be added to the list in the future. +enum DimensionType : byte { + DENSE = 0, + SPARSE_CSR = 1, +} + +table Int32Vector { + values:[int]; +} + +table Uint16Vector { + values:[ushort] (force_align: 4); +} + +table Uint8Vector { + values:[ubyte] (force_align: 4); +} + +// Variable-typed buffer to store the index metadata for a sparse dimension. +// The widest type is Int32 instead of UInt32 because tensor's shape is a int32 +// vector. We don't want the per-dimensional index to overflow that range. +union SparseIndexVector { + Int32Vector, + Uint16Vector, + Uint8Vector +} + +table DimensionMetadata { + // Whether a dimension is dense or sparse. + format:DimensionType; + // Index metadata used for a dimension. + // - If format is DimensionType.DENSE then we use the dense_size field to + // store the size of that dimension. Each index in that dimension is + // stored implicitly. + // - If format is DimensionType.SPARSE_CSR then we use array_segments and + // array_indices to encode that dimension. array_segments represents how + // to segment the indices array, each segment corresponds to one element + // in the previous dimension. array_indices represents the index of the + // non-zero elements within this dimension (as those in the CSR matrix + // format, where the first array is row pointers and the second array is + // column indices). + dense_size:int; + array_segments:SparseIndexVector; + array_indices:SparseIndexVector; +} + +// Parameters to encode a sparse TfLite tensor. +table SparsityParameters { + // The traversal order of the dimensions defined in the `shape` field of the + // conceptual dense tensor. For a n-dimensional tensors with dims (d0, d1, + // ..., dn-1), + // - if not block sparse, the traversal_order is just a permutation of (d0, + // ..., dn-1). For example, a 2-D matrix stored in row-major order would + // have traversal_order = (d0, d1). + // - if block sparse with a k-dimensional block (0 <= k <= n), the + // traversal_order has n + k elements. The first n elements are still a + // permutation of (d0, ..., dn-1). The lask k elements are a permutation + // of (dn, ..., dn+k-1), defining how to traverse a block internally. For + // example, a 2-D matrix with 2-D blocks, both stored in row-major order + // would have traversal_order = (d0, d1, d2, d3). + traversal_order:[int]; + // For an n-dimensional tensor with a k-dimensional block (0 <= k <= n), + // stores how a block dimension in (dn, ..., dn+k-1) maps to the original + // tensor dimension in (d0, ..., dn). + // It's stored in the order of (dn, ..., dn+k-1). + // If not block-sparse, this field is NULL. + block_map:[int]; + // In the traversal order defined above, the metadata needed for + // each dimension to locate the non-zero values in the original dense tensor. + // The size of the dim_metadata array = the size of the traversal_order array + // = n + k. + dim_metadata:[DimensionMetadata]; +} + +table Tensor { + // The tensor shape. The meaning of each entry is operator-specific but + // builtin ops use: [batch size, height, width, number of channels] (That's + // Tensorflow's NHWC). + shape:[int]; + type:TensorType; + // An index that refers to the buffers table at the root of the model. Or, + // if there is no data buffer associated (i.e. intermediate results), then + // this is 0 (which refers to an always existent empty buffer). + // + // The data_buffer itself is an opaque container, with the assumption that the + // target device is little-endian. In addition, all builtin operators assume + // the memory is ordered such that if `shape` is [4, 3, 2], then index + // [i, j, k] maps to data_buffer[i*3*2 + j*2 + k]. + buffer:uint; + name:string; // For debugging and importing back into tensorflow. + quantization:QuantizationParameters; // Optional. + + is_variable:bool = false; + + // Parameters to encode a sparse tensor. See the example in + // tensorflow/lite/testdata/sparse_tensor.json. + sparsity:SparsityParameters; // Optional. + + // Encodes `shape` with unknown dimensions. Unknown dimensions are + // represented with -1. + shape_signature:[int]; // Optional. +} + +// A list of builtin operators. Builtin operators are slightly faster than custom +// ones, but not by much. Moreover, while custom operators accept an opaque +// object containing configuration parameters, builtins have a predetermined +// set of acceptable options. + +enum BuiltinOperator : byte { + ADD = 0, + AVERAGE_POOL_2D = 1, + CONCATENATION = 2, + CONV_2D = 3, + DEPTHWISE_CONV_2D = 4, + DEPTH_TO_SPACE = 5, + DEQUANTIZE = 6, + EMBEDDING_LOOKUP = 7, + FLOOR = 8, + FULLY_CONNECTED = 9, + HASHTABLE_LOOKUP = 10, + L2_NORMALIZATION = 11, + L2_POOL_2D = 12, + LOCAL_RESPONSE_NORMALIZATION = 13, + LOGISTIC = 14, + LSH_PROJECTION = 15, + LSTM = 16, + MAX_POOL_2D = 17, + MUL = 18, + RELU = 19, + // NOTE(aselle): RELU_N1_TO_1 used to be called RELU1, but it was renamed + // since different model developers use RELU1 in different ways. Never + // create another op called RELU1. + RELU_N1_TO_1 = 20, + RELU6 = 21, + RESHAPE = 22, + RESIZE_BILINEAR = 23, + RNN = 24, + SOFTMAX = 25, + SPACE_TO_DEPTH = 26, + SVDF = 27, + TANH = 28, + // Consider rename to CONCATENATE_EMBEDDINGS + CONCAT_EMBEDDINGS = 29, + SKIP_GRAM = 30, + CALL = 31, + CUSTOM = 32, + EMBEDDING_LOOKUP_SPARSE = 33, + PAD = 34, + UNIDIRECTIONAL_SEQUENCE_RNN = 35, + GATHER = 36, + BATCH_TO_SPACE_ND = 37, + SPACE_TO_BATCH_ND = 38, + TRANSPOSE = 39, + MEAN = 40, + SUB = 41, + DIV = 42, + SQUEEZE = 43, + UNIDIRECTIONAL_SEQUENCE_LSTM = 44, + STRIDED_SLICE = 45, + BIDIRECTIONAL_SEQUENCE_RNN = 46, + EXP = 47, + TOPK_V2 = 48, + SPLIT = 49, + LOG_SOFTMAX = 50, + // DELEGATE is a special op type for the operations which are delegated to + // other backends. + // WARNING: Experimental interface, subject to change + DELEGATE = 51, + BIDIRECTIONAL_SEQUENCE_LSTM = 52, + CAST = 53, + PRELU = 54, + MAXIMUM = 55, + ARG_MAX = 56, + MINIMUM = 57, + LESS = 58, + NEG = 59, + PADV2 = 60, + GREATER = 61, + GREATER_EQUAL = 62, + LESS_EQUAL = 63, + SELECT = 64, + SLICE = 65, + SIN = 66, + TRANSPOSE_CONV = 67, + SPARSE_TO_DENSE = 68, + TILE = 69, + EXPAND_DIMS = 70, + EQUAL = 71, + NOT_EQUAL = 72, + LOG = 73, + SUM = 74, + SQRT = 75, + RSQRT = 76, + SHAPE = 77, + POW = 78, + ARG_MIN = 79, + FAKE_QUANT = 80, + REDUCE_PROD = 81, + REDUCE_MAX = 82, + PACK = 83, + LOGICAL_OR = 84, + ONE_HOT = 85, + LOGICAL_AND = 86, + LOGICAL_NOT = 87, + UNPACK = 88, + REDUCE_MIN = 89, + FLOOR_DIV = 90, + REDUCE_ANY = 91, + SQUARE = 92, + ZEROS_LIKE = 93, + FILL = 94, + FLOOR_MOD = 95, + RANGE = 96, + RESIZE_NEAREST_NEIGHBOR = 97, + LEAKY_RELU = 98, + SQUARED_DIFFERENCE = 99, + MIRROR_PAD = 100, + ABS = 101, + SPLIT_V = 102, + UNIQUE = 103, + CEIL = 104, + REVERSE_V2 = 105, + ADD_N = 106, + GATHER_ND = 107, + COS = 108, + WHERE = 109, + RANK = 110, + ELU = 111, + REVERSE_SEQUENCE = 112, + MATRIX_DIAG = 113, + QUANTIZE = 114, + MATRIX_SET_DIAG = 115, + ROUND = 116, + HARD_SWISH = 117, + IF = 118, + WHILE = 119, + NON_MAX_SUPPRESSION_V4 = 120, + NON_MAX_SUPPRESSION_V5 = 121, + SCATTER_ND = 122, + SELECT_V2 = 123, + DENSIFY = 124, + SEGMENT_SUM = 125, + BATCH_MATMUL = 126 +} + + +// Options for the builtin operators. +union BuiltinOptions { + Conv2DOptions, + DepthwiseConv2DOptions, + ConcatEmbeddingsOptions, + LSHProjectionOptions, + Pool2DOptions, + SVDFOptions, + RNNOptions, + FullyConnectedOptions, + SoftmaxOptions, + ConcatenationOptions, + AddOptions, + L2NormOptions, + LocalResponseNormalizationOptions, + LSTMOptions, + ResizeBilinearOptions, + CallOptions, + ReshapeOptions, + SkipGramOptions, + SpaceToDepthOptions, + EmbeddingLookupSparseOptions, + MulOptions, + PadOptions, + GatherOptions, + BatchToSpaceNDOptions, + SpaceToBatchNDOptions, + TransposeOptions, + ReducerOptions, + SubOptions, + DivOptions, + SqueezeOptions, + SequenceRNNOptions, + StridedSliceOptions, + ExpOptions, + TopKV2Options, + SplitOptions, + LogSoftmaxOptions, + CastOptions, + DequantizeOptions, + MaximumMinimumOptions, + ArgMaxOptions, + LessOptions, + NegOptions, + PadV2Options, + GreaterOptions, + GreaterEqualOptions, + LessEqualOptions, + SelectOptions, + SliceOptions, + TransposeConvOptions, + SparseToDenseOptions, + TileOptions, + ExpandDimsOptions, + EqualOptions, + NotEqualOptions, + ShapeOptions, + PowOptions, + ArgMinOptions, + FakeQuantOptions, + PackOptions, + LogicalOrOptions, + OneHotOptions, + LogicalAndOptions, + LogicalNotOptions, + UnpackOptions, + FloorDivOptions, + SquareOptions, + ZerosLikeOptions, + FillOptions, + BidirectionalSequenceLSTMOptions, + BidirectionalSequenceRNNOptions, + UnidirectionalSequenceLSTMOptions, + FloorModOptions, + RangeOptions, + ResizeNearestNeighborOptions, + LeakyReluOptions, + SquaredDifferenceOptions, + MirrorPadOptions, + AbsOptions, + SplitVOptions, + UniqueOptions, + ReverseV2Options, + AddNOptions, + GatherNdOptions, + CosOptions, + WhereOptions, + RankOptions, + ReverseSequenceOptions, + MatrixDiagOptions, + QuantizeOptions, + MatrixSetDiagOptions, + HardSwishOptions, + IfOptions, + WhileOptions, + DepthToSpaceOptions, + NonMaxSuppressionV4Options, + NonMaxSuppressionV5Options, + ScatterNdOptions, + SelectV2Options, + DensifyOptions, + SegmentSumOptions, + BatchMatMulOptions +} + +enum Padding : byte { SAME, VALID } + +enum ActivationFunctionType : byte { + NONE = 0, + RELU = 1, + RELU_N1_TO_1 = 2, + RELU6 = 3, + TANH = 4, + SIGN_BIT = 5, +} + +table Conv2DOptions { + padding:Padding; + stride_w:int; + stride_h:int; + fused_activation_function:ActivationFunctionType; + dilation_w_factor:int = 1; + dilation_h_factor:int = 1; +} + +table Pool2DOptions { + padding:Padding; + stride_w:int; + stride_h:int; + filter_width:int; + filter_height:int; + fused_activation_function:ActivationFunctionType; +} + +table DepthwiseConv2DOptions { + // Parameters for DepthwiseConv version 1 or above. + padding:Padding; + stride_w:int; + stride_h:int; + // `depth_multiplier` is redundant. It's used by CPU kernels in + // TensorFlow 2.0 or below, but ignored in versions above. + // See comments in lite/c/builtin_op_data.h for more details. + depth_multiplier:int; + fused_activation_function:ActivationFunctionType; + // Parameters for DepthwiseConv version 2 or above. + dilation_w_factor:int = 1; + dilation_h_factor:int = 1; +} + +table ConcatEmbeddingsOptions { + num_channels:int; + num_columns_per_channel:[int]; + embedding_dim_per_channel:[int]; // This could be inferred from parameters. +} + +enum LSHProjectionType: byte { + UNKNOWN = 0, + SPARSE = 1, + DENSE = 2, +} + +table LSHProjectionOptions { + type: LSHProjectionType; +} + +table SVDFOptions { + rank:int; + fused_activation_function:ActivationFunctionType; + // For weights-only quantization, use asymmetric quantization for non + // constant inputs at evaluation time. + asymmetric_quantize_inputs:bool; +} + +// An implementation of TensorFlow RNNCell. +table RNNOptions { + fused_activation_function:ActivationFunctionType; + asymmetric_quantize_inputs:bool; +} + +// An implementation of TensorFlow dynamic_rnn with RNNCell. +table SequenceRNNOptions { + time_major:bool; + fused_activation_function:ActivationFunctionType; + asymmetric_quantize_inputs:bool; +} + +// An implementation of TensorFlow bidrectional_dynamic_rnn with RNNCell. +table BidirectionalSequenceRNNOptions { + time_major:bool; + fused_activation_function:ActivationFunctionType; + merge_outputs: bool; + asymmetric_quantize_inputs:bool; +} + +enum FullyConnectedOptionsWeightsFormat: byte { + DEFAULT = 0, + SHUFFLED4x16INT8 = 1, +} + +// An implementation of TensorFlow fully_connected (a.k.a Dense) layer. +table FullyConnectedOptions { + // Parameters for FullyConnected version 1 or above. + fused_activation_function:ActivationFunctionType; + + // Parameters for FullyConnected version 2 or above. + weights_format:FullyConnectedOptionsWeightsFormat = DEFAULT; + + // Parameters for FullyConnected version 5 or above. + // If set to true, then the number of dimension is preserved. Furthermore, + // all but the last dimension of the input and output shapes will be equal. + keep_num_dims: bool; + + // Parameters for FullyConnected version 7 or above. + // If set to true, then weights-only op will use asymmetric quantization for + // inputs. + asymmetric_quantize_inputs: bool; +} + +table SoftmaxOptions { + beta: float; +} + +// An implementation of TensorFlow concat. +table ConcatenationOptions { + axis:int; + fused_activation_function:ActivationFunctionType; +} + +table AddOptions { + fused_activation_function:ActivationFunctionType; +} + +table MulOptions { + fused_activation_function:ActivationFunctionType; +} + +table L2NormOptions { + fused_activation_function:ActivationFunctionType; +} + +table LocalResponseNormalizationOptions { + radius:int; + bias:float; + alpha:float; + beta:float; +} + +enum LSTMKernelType : byte { + // Full LSTM kernel which supports peephole and projection. + FULL = 0, + // Basic LSTM kernels. Equivalent to TensorFlow BasicLSTMCell. + BASIC = 1, +} + +// An implementation of TensorFlow LSTMCell and CoupledInputForgetGateLSTMCell +table LSTMOptions { + // Parameters for LSTM version 1 or above. + fused_activation_function:ActivationFunctionType; + cell_clip: float; // Optional, 0.0 means no clipping + proj_clip: float; // Optional, 0.0 means no clipping + + // Parameters for LSTM version 2 or above. + // Basic kernel is only supported in version 2 or above. + kernel_type: LSTMKernelType = FULL; + + // Parameters for LSTM version 4 or above. + asymmetric_quantize_inputs: bool; +} + +// An implementation of TensorFlow dynamic_rnn with LSTMCell. +table UnidirectionalSequenceLSTMOptions { + fused_activation_function:ActivationFunctionType; + cell_clip: float; // Optional, 0.0 means no clipping + proj_clip: float; // Optional, 0.0 means no clipping + + // If true then first dimension is sequence, otherwise batch. + time_major:bool; + + // Parameter for Unidirectional Sequence LSTM version 4. + asymmetric_quantize_inputs:bool; +} + +table BidirectionalSequenceLSTMOptions { + // Parameters supported by version 1: + fused_activation_function:ActivationFunctionType; + cell_clip: float; // Optional, 0.0 means no clipping + proj_clip: float; // Optional, 0.0 means no clipping + + // If true, store the outputs of both directions into the first output. + merge_outputs: bool; + + // Parameters supported by version 2: + // If true then first dimension is sequence, otherwise batch. + // Version 1 implementations assumed time_major to be true, so this default + // value should never change. + time_major: bool = true; + + // Parameters for version 3 or above. + asymmetric_quantize_inputs:bool; +} + +table ResizeBilinearOptions { + new_height: int (deprecated); + new_width: int (deprecated); + align_corners: bool; + half_pixel_centers: bool; +} + +table ResizeNearestNeighborOptions { + align_corners: bool; + half_pixel_centers: bool; +} + +// A call operation options +table CallOptions { + // The subgraph index that needs to be called. + subgraph:uint; +} + +table PadOptions { +} + +table PadV2Options { +} + +table ReshapeOptions { + new_shape:[int]; +} + +table SpaceToBatchNDOptions { +} + +table BatchToSpaceNDOptions { +} + +table SkipGramOptions { + ngram_size: int; + max_skip_size: int; + include_all_ngrams: bool; +} + +table SpaceToDepthOptions { + block_size: int; +} + +table DepthToSpaceOptions { + block_size: int; +} + +table SubOptions { + fused_activation_function:ActivationFunctionType; +} + +table DivOptions { + fused_activation_function:ActivationFunctionType; +} + +table TopKV2Options { +} + +enum CombinerType : byte { + SUM = 0, + MEAN = 1, + SQRTN = 2, +} + +table EmbeddingLookupSparseOptions { + combiner:CombinerType; +} + +table GatherOptions { + axis: int; +} + +table TransposeOptions { +} + +table ExpOptions { +} + +table CosOptions { +} + +table ReducerOptions { + keep_dims: bool; +} + +table SqueezeOptions { + squeeze_dims:[int]; +} + +table SplitOptions { + num_splits: int; +} + +table SplitVOptions { + num_splits: int; +} + +table StridedSliceOptions { + begin_mask: int; + end_mask: int; + ellipsis_mask: int; + new_axis_mask: int; + shrink_axis_mask: int; +} + +table LogSoftmaxOptions { +} + +table CastOptions { + in_data_type: TensorType; + out_data_type: TensorType; +} + +table DequantizeOptions { +} + +table MaximumMinimumOptions { +} + +table TileOptions { +} + +table ArgMaxOptions { + output_type : TensorType; +} + +table ArgMinOptions { + output_type : TensorType; +} + +table GreaterOptions { +} + +table GreaterEqualOptions { +} + +table LessOptions { +} + +table LessEqualOptions { +} + +table NegOptions { +} + +table SelectOptions { +} + +table SliceOptions { +} + +table TransposeConvOptions { + padding:Padding; + stride_w:int; + stride_h:int; +} + +table ExpandDimsOptions { +} + +table SparseToDenseOptions { + validate_indices:bool; +} + +table EqualOptions { +} + +table NotEqualOptions { +} + +table ShapeOptions { + // Optional output type of the operation (int32 or int64). Defaults to int32. + out_type : TensorType; +} + +table RankOptions { +} + +table PowOptions { +} + +table FakeQuantOptions { + // Parameters supported by version 1: + min:float; + max:float; + num_bits:int; + + // Parameters supported by version 2: + narrow_range:bool; +} + +table PackOptions { + values_count:int; + axis:int; +} + +table LogicalOrOptions { +} + +table OneHotOptions { + axis:int; +} + +table AbsOptions { +} + + +table HardSwishOptions { +} + +table LogicalAndOptions { +} + +table LogicalNotOptions { +} + +table UnpackOptions { + num:int; + axis:int; +} + +table FloorDivOptions { +} + +table SquareOptions { +} + +table ZerosLikeOptions { +} + +table FillOptions { +} + +table FloorModOptions { +} + +table RangeOptions { +} + +table LeakyReluOptions { + alpha:float; +} + +table SquaredDifferenceOptions { +} + +enum MirrorPadMode : byte { + // Doesn't include borders. + REFLECT = 0, + // Includes borders. + SYMMETRIC = 1, +} + +table MirrorPadOptions { + mode:MirrorPadMode; +} + +table UniqueOptions { + idx_out_type:TensorType = INT32; +} + +table ReverseV2Options { +} + +table AddNOptions { +} + +table GatherNdOptions { +} + +table WhereOptions { +} + +table ReverseSequenceOptions { + seq_dim:int; + batch_dim:int = 0; +} + +table MatrixDiagOptions { +} + +table QuantizeOptions { +} + +table MatrixSetDiagOptions { +} + +table IfOptions { + then_subgraph_index:int; + else_subgraph_index:int; +} + +table WhileOptions { + cond_subgraph_index:int; + body_subgraph_index:int; +} + +table NonMaxSuppressionV4Options { +} + +table NonMaxSuppressionV5Options { +} + +table ScatterNdOptions { +} + +table SelectV2Options { +} + +table DensifyOptions { +} + +table SegmentSumOptions { +} + +table BatchMatMulOptions { + adj_x:bool; + adj_y:bool; +} + +// An OperatorCode can be an enum value (BuiltinOperator) if the operator is a +// builtin, or a string if the operator is custom. +table OperatorCode { + builtin_code:BuiltinOperator; + custom_code:string; + + // The version of the operator. The version need to be bumped whenever new + // parameters are introduced into an op. + version:int = 1; +} + +enum CustomOptionsFormat : byte { + FLEXBUFFERS = 0, +} + +// An operator takes tensors as inputs and outputs. The type of operation being +// performed is determined by an index into the list of valid OperatorCodes, +// while the specifics of each operations is configured using builtin_options +// or custom_options. +table Operator { + // Index into the operator_codes array. Using an integer here avoids + // complicate map lookups. + opcode_index:uint; + + // Optional input are indicated by -1. + inputs:[int]; + outputs:[int]; + + builtin_options:BuiltinOptions; + custom_options:[ubyte]; + custom_options_format:CustomOptionsFormat; + + // A list of booleans indicating the input tensors which are being mutated by + // this operator.(e.g. used by RNN and LSTM). + // For example, if the "inputs" array refers to 5 tensors and the second and + // fifth are mutable variables, then this list will contain + // [false, true, false, false, true]. + // + // If the list is empty, no variable is mutated in this operator. + // The list either has the same length as `inputs`, or is empty. + mutating_variable_inputs:[bool]; + + // A list of indices to the subgraph's "tensors" that are internal to an Op. + // Internal tensors are those that do not flow in or out of the operation, + // but instead are part of internal computation. As such, the operation's + // implementation may manage its memory more efficiently. They are needed + // however (i.e. not just an implementation detail) since they are part of the + // computation, which may require relevant metadata such as quantization + // parameters. + intermediates:[int]; +} + +// The root type, defining a subgraph, which typically represents an entire +// model. +table SubGraph { + // A list of all tensors used in this subgraph. + tensors:[Tensor]; + + // Indices of the tensors that are inputs into this subgraph. Note this is + // the list of non-static tensors that feed into the subgraph for inference. + inputs:[int]; + + // Indices of the tensors that are outputs out of this subgraph. Note this is + // the list of output tensors that are considered the product of the + // subgraph's inference. + outputs:[int]; + + // All operators, in execution order. + operators:[Operator]; + + // Name of this subgraph (used for debugging). + name:string; +} + +// Table of raw data buffers (used for constant tensors). Referenced by tensors +// by index. The generous alignment accommodates mmap-friendly data structures. +table Buffer { + data:[ubyte] (force_align: 16); +} + +table Metadata { + // A human readable string to uniquely identify a Metadata. + name:string; + // An index to the buffers table. + buffer:uint; +} + +table Model { + // Version of the schema. + version:uint; + + // A list of all operator codes used in this model. This is + // kept in order because operators carry an index into this + // vector. + operator_codes:[OperatorCode]; + + // All the subgraphs of the model. The 0th is assumed to be the main + // model. + subgraphs:[SubGraph]; + + // A description of the model. + description:string; + + // Buffers of the model. + // Note the 0th entry of this array must be an empty buffer (sentinel). + // This is a convention so that tensors without a buffer can provide 0 as + // their buffer. + buffers:[Buffer]; + + // Metadata about the model. Indirects into the existings buffers list. + // Deprecated, prefer to use metadata field. + metadata_buffer:[int]; + + // Metadata about the model. + metadata:[Metadata]; +} + +root_type Model; diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.cc index 2a39a94747..3d33607634 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.cc @@ -24,33 +24,33 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteReluParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteReluParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::RELU); prim->set_min_val(0); prim->set_max_val(FLT_MAX); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteRelu6Parser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteRelu6Parser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::RELU6); prim->set_min_val(0); prim->set_max_val(kValueThreshold6); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteLeakyReluParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLeakyReluParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::LEAKY_RELU); @@ -59,47 +59,47 @@ ops::PrimitiveC *TfliteLeakyReluParser::Parse(const std::unique_ptrset_alpha(tflite_attr->alpha); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TflitePReLUParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TflitePReLUParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_channel_shared(true); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteTanhParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteTanhParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::TANH); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteHardSwishParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteHardSwishParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::HSWISH); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteLogisticParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLogisticParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_activation_type(mindspore::ActivationType::SIGMOID); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_TfliteReluParser(tflite::BuiltinOperator_RELU, new TfliteReluParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.h index c82c1a0c23..ec13e59fad 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_activation_parser.h @@ -16,6 +16,7 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_PARSER_TFLITE_ACTIVATION_PARSER_H #define MINDSPORE_LITE_TOOLS_CONVERTER_PARSER_TFLITE_ACTIVATION_PARSER_H +#define USE_DEPRECATED_API #include #include @@ -31,9 +32,9 @@ class TfliteReluParser : public TfliteNodeParser { ~TfliteReluParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteRelu6Parser : public TfliteNodeParser { @@ -42,9 +43,9 @@ class TfliteRelu6Parser : public TfliteNodeParser { ~TfliteRelu6Parser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLeakyReluParser : public TfliteNodeParser { @@ -53,9 +54,9 @@ class TfliteLeakyReluParser : public TfliteNodeParser { ~TfliteLeakyReluParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TflitePReLUParser : public TfliteNodeParser { @@ -64,9 +65,9 @@ class TflitePReLUParser : public TfliteNodeParser { ~TflitePReLUParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteTanhParser : public TfliteNodeParser { @@ -75,9 +76,9 @@ class TfliteTanhParser : public TfliteNodeParser { ~TfliteTanhParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteHardSwishParser : public TfliteNodeParser { @@ -86,9 +87,9 @@ class TfliteHardSwishParser : public TfliteNodeParser { ~TfliteHardSwishParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLogisticParser : public TfliteNodeParser { @@ -97,9 +98,9 @@ class TfliteLogisticParser : public TfliteNodeParser { ~TfliteLogisticParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.cc index 29c2dced60..2d94ec22be 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.cc @@ -23,12 +23,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteAddNParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteAddNParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteAddNParser(tflite::BuiltinOperator_ADD_N, new TfliteAddNParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.h index e84cdc2955..63ca61f3af 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_addn_parser.h @@ -31,9 +31,9 @@ class TfliteAddNParser : public TfliteNodeParser { ~TfliteAddNParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.cc index 1dbc0ff81a..1dfe52bbcd 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteArgmaxParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteArgmaxParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_keep_dims(false); @@ -43,7 +43,7 @@ ops::PrimitiveC *TfliteArgmaxParser::Parse(const std::unique_ptrset_axis(axes.at(0)); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteArgmaxParser(tflite::BuiltinOperator_ARG_MAX, new TfliteArgmaxParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.h index 44f30f6d65..79ad4e3364 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_argmax_parser.h @@ -31,9 +31,9 @@ class TfliteArgmaxParser : public TfliteNodeParser { ~TfliteArgmaxParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.cc index e0faae5e4a..6d8876fed8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteArgminParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteArgminParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_keep_dims(false); @@ -43,7 +43,7 @@ ops::PrimitiveC *TfliteArgminParser::Parse(const std::unique_ptrset_axis(axes.at(0)); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteArgminParser(tflite::BuiltinOperator_ARG_MIN, new TfliteArgminParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.h index d4d4c32e38..51f44ea1f9 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_argmin_parser.h @@ -31,9 +31,9 @@ class TfliteArgminParser : public TfliteNodeParser { ~TfliteArgminParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.cc index 5da69cb8f8..962bffbd92 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.cc @@ -49,9 +49,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteAddParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteAddParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -59,12 +59,12 @@ ops::PrimitiveC *TfliteAddParser::Parse(const std::unique_ptr MS_CHECK_TRUE_MSG(tflite_attr != nullptr, nullptr, "get AddFusion attr failed"); prim->set_activation_type(GetActivationFunctionType(tflite_attr->fused_activation_function)); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteMulParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteMulParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); const auto &tflite_attr = tflite_op->builtin_options.AsMulOptions(); @@ -74,12 +74,12 @@ ops::PrimitiveC *TfliteMulParser::Parse(const std::unique_ptr } prim->set_activation_type(GetActivationFunctionType(tflite_attr->fused_activation_function)); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteDivParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteDivParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); const auto &tflite_attr = tflite_op->builtin_options.AsDivOptions(); @@ -89,12 +89,12 @@ ops::PrimitiveC *TfliteDivParser::Parse(const std::unique_ptr } prim->set_activation_type(GetActivationFunctionType(tflite_attr->fused_activation_function)); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteSubParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSubParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); const auto &tflite_attr = tflite_op->builtin_options.AsSubOptions(); @@ -104,87 +104,87 @@ ops::PrimitiveC *TfliteSubParser::Parse(const std::unique_ptr } prim->set_activation_type(GetActivationFunctionType(tflite_attr->fused_activation_function)); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteFloorDivParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteFloorDivParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteFloorModParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteFloorModParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TflitePowParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TflitePowParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_scale(1.0); prim->set_shift(0.0); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteSquaredDifferenceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSquaredDifferenceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteMaximumParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteMaximumParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteMinimumParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteMinimumParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteAbsParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteAbsParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCosParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteCosParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteFloorParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteFloorParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteExpParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteExpParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -192,119 +192,119 @@ ops::PrimitiveC *TfliteExpParser::Parse(const std::unique_ptr prim->set_scale(1.0); prim->set_shift(0.0); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCeilParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteCeilParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteLogParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLogParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteRoundParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteRoundParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteSqrtParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSqrtParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteRsqrtParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteRsqrtParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteSquareParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSquareParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteSinParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSinParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteNegParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteNegParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteEqualParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteEqualParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteNotEqualParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteNotEqualParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteGreaterParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { - auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); -} - -ops::PrimitiveC *TfliteGreaterEqualParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { - auto prim = std::make_unique(); - MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); -} - -ops::PrimitiveC *TfliteLessParser::Parse(const std::unique_ptr &tflite_op, +PrimitiveCPtr TfliteGreaterParser::Parse(const std::unique_ptr &tflite_op, const std::unique_ptr &tflite_subgraph, const std::unique_ptr &tflite_model) { - auto prim = std::make_unique(); + auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteLessEqualParser::Parse(const std::unique_ptr &tflite_op, +PrimitiveCPtr TfliteGreaterEqualParser::Parse(const std::unique_ptr &tflite_op, const std::unique_ptr &tflite_subgraph, const std::unique_ptr &tflite_model) { + auto prim = std::make_unique(); + MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + return prim->GetPrim(); +} + +PrimitiveCPtr TfliteLessParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { + auto prim = std::make_unique(); + MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + return prim->GetPrim(); +} + +PrimitiveCPtr TfliteLessEqualParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteAddParser(tflite::BuiltinOperator_ADD, new TfliteAddParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.h index 90ae85e2ca..705f8b8c88 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_arithmetic_parser.h @@ -31,9 +31,9 @@ class TfliteAddParser : public TfliteNodeParser { ~TfliteAddParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteSubParser : public TfliteNodeParser { @@ -42,9 +42,9 @@ class TfliteSubParser : public TfliteNodeParser { ~TfliteSubParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteMulParser : public TfliteNodeParser { @@ -53,9 +53,9 @@ class TfliteMulParser : public TfliteNodeParser { ~TfliteMulParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteDivParser : public TfliteNodeParser { @@ -64,9 +64,9 @@ class TfliteDivParser : public TfliteNodeParser { ~TfliteDivParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteFloorDivParser : public TfliteNodeParser { @@ -75,9 +75,9 @@ class TfliteFloorDivParser : public TfliteNodeParser { ~TfliteFloorDivParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteFloorModParser : public TfliteNodeParser { @@ -86,9 +86,9 @@ class TfliteFloorModParser : public TfliteNodeParser { ~TfliteFloorModParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TflitePowParser : public TfliteNodeParser { @@ -97,9 +97,9 @@ class TflitePowParser : public TfliteNodeParser { ~TflitePowParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteSquaredDifferenceParser : public TfliteNodeParser { @@ -108,9 +108,9 @@ class TfliteSquaredDifferenceParser : public TfliteNodeParser { ~TfliteSquaredDifferenceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteMaximumParser : public TfliteNodeParser { @@ -119,9 +119,9 @@ class TfliteMaximumParser : public TfliteNodeParser { ~TfliteMaximumParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteMinimumParser : public TfliteNodeParser { @@ -130,9 +130,9 @@ class TfliteMinimumParser : public TfliteNodeParser { ~TfliteMinimumParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteAbsParser : public TfliteNodeParser { @@ -141,9 +141,9 @@ class TfliteAbsParser : public TfliteNodeParser { ~TfliteAbsParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteExpParser : public TfliteNodeParser { @@ -152,9 +152,9 @@ class TfliteExpParser : public TfliteNodeParser { ~TfliteExpParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteSqrtParser : public TfliteNodeParser { @@ -163,9 +163,9 @@ class TfliteSqrtParser : public TfliteNodeParser { ~TfliteSqrtParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteRsqrtParser : public TfliteNodeParser { @@ -174,9 +174,9 @@ class TfliteRsqrtParser : public TfliteNodeParser { ~TfliteRsqrtParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteSquareParser : public TfliteNodeParser { @@ -185,9 +185,9 @@ class TfliteSquareParser : public TfliteNodeParser { ~TfliteSquareParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteSinParser : public TfliteNodeParser { @@ -196,9 +196,9 @@ class TfliteSinParser : public TfliteNodeParser { ~TfliteSinParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteCosParser : public TfliteNodeParser { @@ -207,9 +207,9 @@ class TfliteCosParser : public TfliteNodeParser { ~TfliteCosParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLogParser : public TfliteNodeParser { @@ -218,9 +218,9 @@ class TfliteLogParser : public TfliteNodeParser { ~TfliteLogParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteRoundParser : public TfliteNodeParser { @@ -229,9 +229,9 @@ class TfliteRoundParser : public TfliteNodeParser { ~TfliteRoundParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteCeilParser : public TfliteNodeParser { @@ -240,9 +240,9 @@ class TfliteCeilParser : public TfliteNodeParser { ~TfliteCeilParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteFloorParser : public TfliteNodeParser { @@ -251,9 +251,9 @@ class TfliteFloorParser : public TfliteNodeParser { ~TfliteFloorParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteNegParser : public TfliteNodeParser { @@ -262,9 +262,9 @@ class TfliteNegParser : public TfliteNodeParser { ~TfliteNegParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteEqualParser : public TfliteNodeParser { @@ -273,9 +273,9 @@ class TfliteEqualParser : public TfliteNodeParser { ~TfliteEqualParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteNotEqualParser : public TfliteNodeParser { @@ -284,9 +284,9 @@ class TfliteNotEqualParser : public TfliteNodeParser { ~TfliteNotEqualParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteGreaterParser : public TfliteNodeParser { @@ -295,9 +295,9 @@ class TfliteGreaterParser : public TfliteNodeParser { ~TfliteGreaterParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteGreaterEqualParser : public TfliteNodeParser { @@ -306,9 +306,9 @@ class TfliteGreaterEqualParser : public TfliteNodeParser { ~TfliteGreaterEqualParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLessParser : public TfliteNodeParser { @@ -317,9 +317,9 @@ class TfliteLessParser : public TfliteNodeParser { ~TfliteLessParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLessEqualParser : public TfliteNodeParser { @@ -328,9 +328,9 @@ class TfliteLessEqualParser : public TfliteNodeParser { ~TfliteLessEqualParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.cc index a1322e07b5..febb485861 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteBatchMatmulParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteBatchMatmulParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -36,7 +36,7 @@ ops::PrimitiveC *TfliteBatchMatmulParser::Parse(const std::unique_ptrset_transpose_a(tflite_attr->adj_x); prim->set_transpose_b(tflite_attr->adj_y); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteBatchMatmulParser(tflite::BuiltinOperator_BATCH_MATMUL, new TfliteBatchMatmulParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.h index 9014a6da59..ade86dd868 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_matmul_parser.h @@ -31,9 +31,9 @@ class TfliteBatchMatmulParser : public TfliteNodeParser { ~TfliteBatchMatmulParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.cc index 566bdb2d38..03886d4d94 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.cc @@ -24,9 +24,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteBatchToSpaceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteBatchToSpaceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize2, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -46,7 +46,7 @@ ops::PrimitiveC *TfliteBatchToSpaceParser::Parse(const std::unique_ptrset_crops(crops); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteBatchToSpaceNDParser(tflite::BuiltinOperator_BATCH_TO_SPACE_ND, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.h index ca48356839..ff5c3ea466 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_batch_to_space_parser.h @@ -31,9 +31,9 @@ class TfliteBatchToSpaceParser : public TfliteNodeParser { ~TfliteBatchToSpaceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.cc index a3dcc07ec1..2fc5ed0b88 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteBroadcastToParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteBroadcastToParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize1, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -37,7 +37,7 @@ ops::PrimitiveC *TfliteBroadcastToParser::Parse(const std::unique_ptrset_shape(dst_shape); - return prim.release(); + return prim->GetPrim(); } } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.h index ed8c555b9f..676568f64d 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_broadcast_to_parser.h @@ -31,9 +31,9 @@ class TfliteBroadcastToParser : public TfliteNodeParser { ~TfliteBroadcastToParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.cc index 9d753b5c67..0a75e2904a 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.cc @@ -22,12 +22,14 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteCastParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteCastParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_TRUE_RET(!tflite_op->outputs.empty(), nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); const auto &out_tensor = tflite_subgraph->tensors[tflite_op->outputs.front()]; if (out_tensor == nullptr) { @@ -37,9 +39,9 @@ ops::PrimitiveC *TfliteCastParser::Parse(const std::unique_ptrtype); auto value_dst = MakeValue(static_cast(dstT)); MS_CHECK_TRUE_RET(value_dst != nullptr, nullptr); - prim->AddAttr("to", value_dst); + prim_c->AddAttr("to", value_dst); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteCastParser(tflite::BuiltinOperator_CAST, new TfliteCastParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.h index add54035f2..b0d70a5b2b 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_cast_parser.h @@ -31,9 +31,9 @@ class TfliteCastParser : public TfliteNodeParser { ~TfliteCastParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.cc index f6c5639a3e..ffcc3f88f1 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteConcatParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteConcatParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -35,7 +35,7 @@ ops::PrimitiveC *TfliteConcatParser::Parse(const std::unique_ptrset_axis(tflite_attr->axis); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteConcatParser(tflite::BuiltinOperator_CONCATENATION, new TfliteConcatParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.h index f11c44ea66..2d7420b822 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_concat_parser.h @@ -31,9 +31,9 @@ class TfliteConcatParser : public TfliteNodeParser { ~TfliteConcatParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.cc index f626286a1e..34c0b136f8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.cc @@ -92,9 +92,9 @@ STATUS GetConvPaddingParam(const std::unique_ptr &tensor, minds return RET_OK; } } // namespace -ops::PrimitiveC *TfliteConvParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteConvParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -144,12 +144,12 @@ ops::PrimitiveC *TfliteConvParser::Parse(const std::unique_ptrset_pad_list(params); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteDepthwiseConv2DParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteDepthwiseConv2DParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -218,9 +218,11 @@ ops::PrimitiveC *TfliteDepthwiseConv2DParser::Parse(const std::unique_ptr(true); MS_CHECK_TRUE_RET(value_ptr != nullptr, nullptr); - prim->AddAttr(ops::kIsDepthWise, value_ptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); + prim_c->AddAttr(ops::kIsDepthWise, value_ptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteConv2DParser(tflite::BuiltinOperator_CONV_2D, new TfliteConvParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.h index eb433442df..245c3fb035 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_parser.h @@ -31,9 +31,9 @@ class TfliteConvParser : public TfliteNodeParser { ~TfliteConvParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteDepthwiseConv2DParser : public TfliteNodeParser { @@ -42,9 +42,9 @@ class TfliteDepthwiseConv2DParser : public TfliteNodeParser { ~TfliteDepthwiseConv2DParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.cc index 70d0ac673d..8f943fa329 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteDeConvParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteDeConvParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -75,7 +75,7 @@ ops::PrimitiveC *TfliteDeConvParser::Parse(const std::unique_ptrset_pad_list(params); } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteDeConv2DParser(tflite::BuiltinOperator_TRANSPOSE_CONV, new TfliteDeConvParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.h index a1b7e7e1be..e862c652f3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_conv_transpose_parser.h @@ -31,9 +31,9 @@ class TfliteDeConvParser : public TfliteNodeParser { ~TfliteDeConvParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.cc index 782a1bffeb..1588eca28d 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.cc @@ -34,9 +34,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteCustomParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteCustomParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { const auto &custom_attr = tflite_op->custom_options; const auto &opnode = tflite_model->operator_codes.at(tflite_op->opcode_index); if (opnode == nullptr) { @@ -68,8 +68,8 @@ ops::PrimitiveC *TfliteCustomParser::Parse(const std::unique_ptr &custom_attr, - const std::unique_ptr &tflite_op) { +PrimitiveCPtr TfliteCustomParser::DetectPostProcess(const std::vector &custom_attr, + const std::unique_ptr &tflite_op) { MS_CHECK_TRUE_RET(tflite_op != nullptr, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -101,10 +101,10 @@ ops::PrimitiveC *TfliteCustomParser::DetectPostProcess(const std::vectorset_out_quantized(attr_map["_output_quantized"].AsBool()); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::AudioSpectrogram(const std::vector &custom_attr) { +PrimitiveCPtr TfliteCustomParser::AudioSpectrogram(const std::vector &custom_attr) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -113,10 +113,10 @@ ops::PrimitiveC *TfliteCustomParser::AudioSpectrogram(const std::vector prim->set_stride(attr_map["stride"].AsInt64()); prim->set_mag_square(attr_map["magnitude_squared"].AsBool()); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::Mfcc(const std::vector &custom_attr) { +PrimitiveCPtr TfliteCustomParser::Mfcc(const std::vector &custom_attr) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -126,10 +126,10 @@ ops::PrimitiveC *TfliteCustomParser::Mfcc(const std::vector &custom_att prim->set_filter_bank_channel_num(attr_map["filterbank_channel_count"].AsInt64()); prim->set_dct_coeff_num(attr_map["dct_coefficient_count"].AsInt64()); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::Predict(const std::vector &custom_attr) { +PrimitiveCPtr TfliteCustomParser::Predict(const std::vector &custom_attr) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); MS_CHECK_TRUE_RET(custom_attr.data() != nullptr, nullptr); @@ -138,25 +138,25 @@ ops::PrimitiveC *TfliteCustomParser::Predict(const std::vector &custom_ auto weight_thres = reinterpret_cast(custom_attr.data())[1]; prim->set_output_num(static_cast(out_num)); prim->set_weight_threshold(weight_thres); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::Normalize() { +PrimitiveCPtr TfliteCustomParser::Normalize() { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::ExtractFeatures() { +PrimitiveCPtr TfliteCustomParser::ExtractFeatures() { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::Rfft(const std::vector &custom_attr, - const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteCustomParser::Rfft(const std::vector &custom_attr, + const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_TRUE_RET(tflite_op != nullptr, nullptr); MS_CHECK_TRUE_RET(tflite_subgraph != nullptr, nullptr); MS_CHECK_TRUE_RET(tflite_model != nullptr, nullptr); @@ -171,25 +171,25 @@ ops::PrimitiveC *TfliteCustomParser::Rfft(const std::vector &custom_att } prim->set_fft_length(fft_length[0]); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::FftReal() { +PrimitiveCPtr TfliteCustomParser::FftReal() { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::FftImag() { +PrimitiveCPtr TfliteCustomParser::FftImag() { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteCustomParser::Identity() { +PrimitiveCPtr TfliteCustomParser::Identity() { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteCustomParser(tflite::BuiltinOperator_CUSTOM, new TfliteCustomParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.h index 1294bb78e6..8c4e078381 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_custom_parser.h @@ -31,32 +31,32 @@ class TfliteCustomParser : public TfliteNodeParser { ~TfliteCustomParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; - static ops::PrimitiveC *DetectPostProcess(const std::vector &custom_attr, - const std::unique_ptr &tflite_op); + static PrimitiveCPtr DetectPostProcess(const std::vector &custom_attr, + const std::unique_ptr &tflite_op); - static ops::PrimitiveC *AudioSpectrogram(const std::vector &custom_attr); + static PrimitiveCPtr AudioSpectrogram(const std::vector &custom_attr); - static ops::PrimitiveC *Mfcc(const std::vector &custom_attr); + static PrimitiveCPtr Mfcc(const std::vector &custom_attr); - static ops::PrimitiveC *Predict(const std::vector &custom_attr); + static PrimitiveCPtr Predict(const std::vector &custom_attr); - static ops::PrimitiveC *Normalize(); + static PrimitiveCPtr Normalize(); - static ops::PrimitiveC *ExtractFeatures(); + static PrimitiveCPtr ExtractFeatures(); - ops::PrimitiveC *Rfft(const std::vector &custom_attr, const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model); + PrimitiveCPtr Rfft(const std::vector &custom_attr, const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model); - static ops::PrimitiveC *FftReal(); + static PrimitiveCPtr FftReal(); - static ops::PrimitiveC *FftImag(); + static PrimitiveCPtr FftImag(); - static ops::PrimitiveC *Identity(); + static PrimitiveCPtr Identity(); }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.cc index 523f7926e6..b76ff3f232 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteDepthToSpaceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteDepthToSpaceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -38,7 +38,7 @@ ops::PrimitiveC *TfliteDepthToSpaceParser::Parse(const std::unique_ptrset_block_size(tflite_attr->block_size); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteDepthToSpaceParser(tflite::BuiltinOperator_DEPTH_TO_SPACE, new TfliteDepthToSpaceParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.h index 86cf3652c1..933446718c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_depth_to_space_parser.h @@ -31,9 +31,9 @@ class TfliteDepthToSpaceParser : public TfliteNodeParser { ~TfliteDepthToSpaceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.cc index 7c1b45196a..b8d53ceb3c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteDequantizeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteDequantizeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_TRUE_RET(!tflite_op->inputs.empty(), nullptr); MS_CHECK_TRUE_RET(!tflite_op->outputs.empty(), nullptr); const auto &in_tensor = tflite_subgraph->tensors[tflite_op->inputs.at(FIRST_INPUT)]; @@ -43,15 +43,17 @@ ops::PrimitiveC *TfliteDequantizeParser::Parse(const std::unique_ptrset_src_t(GetTfliteDataType(in_tensor->type)); prim->set_dst_t(GetTfliteDataType(out_tensor->type)); - return prim.release(); + return prim->GetPrim(); } else { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); auto dstT = GetTfliteDataType(out_tensor->type); auto value_dst = MakeValue(static_cast(dstT)); MS_CHECK_TRUE_RET(value_dst != nullptr, nullptr); - prim->AddAttr("to", value_dst); - return prim.release(); + prim_c->AddAttr("to", value_dst); + return prim->GetPrim(); } } diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.h index 91d0b05cf1..5b988f6cf7 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_dequantize_parser.h @@ -30,9 +30,9 @@ class TfliteDequantizeParser : public TfliteNodeParser { ~TfliteDequantizeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.cc index cc96b1e359..f6c0a42b14 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteExpandDimsParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteExpandDimsParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteExpandDimsParser(tflite::BuiltinOperator_EXPAND_DIMS, new TfliteExpandDimsParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.h index 684394fb62..3ab7185add 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_expand_dims_parser.h @@ -31,9 +31,9 @@ class TfliteExpandDimsParser : public TfliteNodeParser { ~TfliteExpandDimsParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.cc index 3f6446e707..0c18bacbd3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteFillParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteFillParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteFillParser(tflite::BuiltinOperator_FILL, new TfliteFillParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.h index 4e45d76b97..e9e47ea526 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_fill_parser.h @@ -31,9 +31,9 @@ class TfliteFillParser : public TfliteNodeParser { ~TfliteFillParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.cc index 2c36cca28c..c16587c8e3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteFullyConnectedParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteFullyConnectedParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -39,7 +39,7 @@ ops::PrimitiveC *TfliteFullyConnectedParser::Parse(const std::unique_ptrset_activation_type(GetActivationFunctionType(tflite_attr->fused_activation_function)); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteFullyConnectedParser(tflite::BuiltinOperator_FULLY_CONNECTED, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.h index 1aeb66ca9a..a5c2ca0c3a 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_fullyconnected_parser.h @@ -31,9 +31,9 @@ class TfliteFullyConnectedParser : public TfliteNodeParser { ~TfliteFullyConnectedParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.cc index 0f41733d5d..60568dc02a 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteGatherNdParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteGatherNdParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteGatherNdParser(tflite::BuiltinOperator_GATHER_ND, new TfliteGatherNdParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.h index 1a0cd71fce..50a495424a 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_nd_parser.h @@ -31,9 +31,9 @@ class TfliteGatherNdParser : public TfliteNodeParser { ~TfliteGatherNdParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.cc index 70a7735e2c..6a99a454ac 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.cc @@ -22,11 +22,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteGatherParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteGatherParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); const auto &tflite_attr = tflite_op->builtin_options.AsGatherOptions(); if (tflite_attr == nullptr) { @@ -35,9 +37,9 @@ ops::PrimitiveC *TfliteGatherParser::Parse(const std::unique_ptr(tflite_attr->axis)); MS_CHECK_TRUE_RET(axis_value != nullptr, nullptr); - prim->AddAttr("axis", axis_value); + prim_c->AddAttr("axis", axis_value); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteGatherParser(tflite::BuiltinOperator_GATHER, new TfliteGatherParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.h index 4767ce1f40..d73d822adf 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_gather_parser.h @@ -31,9 +31,9 @@ class TfliteGatherParser : public TfliteNodeParser { ~TfliteGatherParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.cc index c077d3659c..f8a24a96fa 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteHashtableLookupParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteHashtableLookupParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteHashtableLookupParser(tflite::BuiltinOperator_HASHTABLE_LOOKUP, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.h index d869fba719..90a3096235 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_hashtable_lookup_parser.h @@ -31,9 +31,9 @@ class TfliteHashtableLookupParser : public TfliteNodeParser { ~TfliteHashtableLookupParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.cc index 2c498d1de1..65bb608e7b 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.cc @@ -22,13 +22,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteIfParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { - auto prim = std::make_unique(); +PrimitiveCPtr TfliteIfParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim; } TfliteNodeRegister g_tfliteIfParser(tflite::BuiltinOperator_IF, new TfliteIfParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.h index 93c1d5524c..9bfb26f84c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_if_parser.h @@ -32,9 +32,9 @@ class TfliteIfParser : public TfliteNodeParser { ~TfliteIfParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_inputs_adjust.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_inputs_adjust.cc index f9d512966c..a2c697acf4 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_inputs_adjust.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_inputs_adjust.cc @@ -40,7 +40,7 @@ bool CheckResize(const CNodePtr &cnode) { if (!opt::CheckPrimitiveType(cnode, prim::kPrimResize)) { return false; } - auto prim_resize = GetValueNode>(cnode->input(0)); + auto prim_resize = ops::GetOperator(cnode->input(0)); if (prim_resize == nullptr || prim_resize->GetAttr(ops::kNewHeight) == nullptr || prim_resize->GetAttr(ops::kNewWidth) == nullptr) { return false; diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.cc index d43112809a..fb139b64d5 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteL2NormParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteL2NormParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -39,7 +39,7 @@ ops::PrimitiveC *TfliteL2NormParser::Parse(const std::unique_ptrset_activation_type(GetActivationFunctionType(tflite_attr->fused_activation_function)); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteL2NormParser(tflite::BuiltinOperator_L2_NORMALIZATION, new TfliteL2NormParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.h index 1ce8e799bb..c70f2dde65 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_l2norm_parser.h @@ -31,9 +31,9 @@ class TfliteL2NormParser : public TfliteNodeParser { ~TfliteL2NormParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.cc index 163e532937..58a2e091a9 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.cc @@ -22,15 +22,15 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteLogSoftmaxParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLogSoftmaxParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_axis(-1); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteLogSoftmaxParser(tflite::BuiltinOperator_LOG_SOFTMAX, new TfliteLogSoftmaxParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.h index 37bf4fe532..c86b7c786c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_log_softmax_parser.h @@ -31,9 +31,9 @@ class TfliteLogSoftmaxParser : public TfliteNodeParser { ~TfliteLogSoftmaxParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.cc index 0b4d32ddc7..c7de7818fb 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.cc @@ -24,28 +24,28 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteLogicalAndParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLogicalAndParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteLogicalNotParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLogicalNotParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteLogicalOrParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLogicalOrParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteLogicalAndParser(tflite::BuiltinOperator_LOGICAL_AND, new TfliteLogicalAndParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.h index 5b47546dfa..b4a6dc85bf 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_logical_parser.h @@ -31,9 +31,9 @@ class TfliteLogicalAndParser : public TfliteNodeParser { ~TfliteLogicalAndParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLogicalNotParser : public TfliteNodeParser { @@ -42,9 +42,9 @@ class TfliteLogicalNotParser : public TfliteNodeParser { ~TfliteLogicalNotParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteLogicalOrParser : public TfliteNodeParser { @@ -53,9 +53,9 @@ class TfliteLogicalOrParser : public TfliteNodeParser { ~TfliteLogicalOrParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.cc index d9669fa28e..0bb721de84 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteLRNParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLRNParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -38,7 +38,7 @@ ops::PrimitiveC *TfliteLRNParser::Parse(const std::unique_ptr prim->set_beta(tflite_attr->beta); prim->set_bias(tflite_attr->bias); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteLRNParser(tflite::BuiltinOperator_LOCAL_RESPONSE_NORMALIZATION, new TfliteLRNParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.h index 71721042c4..d4644e9792 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_lrn_parser.h @@ -31,9 +31,9 @@ class TfliteLRNParser : public TfliteNodeParser { ~TfliteLRNParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.cc index 8b38b949f5..ff4c0f50d6 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteLshProjectionParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteLshProjectionParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -44,7 +44,7 @@ ops::PrimitiveC *TfliteLshProjectionParser::Parse(const std::unique_ptrset_type(mindspore::LshProjectionType::UNKNOWN); } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteLshProjectionParser(tflite::BuiltinOperator_LSH_PROJECTION, new TfliteLshProjectionParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.h index 35385b523b..fd6296af02 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_lsh_projection_parser.h @@ -31,9 +31,9 @@ class TfliteLshProjectionParser : public TfliteNodeParser { ~TfliteLshProjectionParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.cc index 9a24f02344..4c0a72f4b8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.cc @@ -21,9 +21,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteMatMulParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteMatMulParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -36,7 +36,7 @@ ops::PrimitiveC *TfliteMatMulParser::Parse(const std::unique_ptrset_transpose_b(tflite_attr->adj_y); prim->set_activation_type(mindspore::NO_ACTIVATION); - return prim.release(); + return prim->GetPrim(); } } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.h index 0c74c13639..14531191c4 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_matmul_parser.h @@ -30,9 +30,9 @@ class TfliteMatMulParser : public TfliteNodeParser { ~TfliteMatMulParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.cc index 0cc4aa9ecd..745ed28ed6 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.cc @@ -42,6 +42,11 @@ namespace mindspore::lite { namespace { constexpr size_t kMainGraphIndex = 0; constexpr size_t kConvWeightIndex = 2; + +FuncGraphPtr ConvertGraph(const api::FuncGraphPtr &func_graph) { + auto impl = func_graph->impl(); + return std::dynamic_pointer_cast(impl); +} } // namespace std::unique_ptr TfliteModelParser::ReadTfliteModel(const std::string &model_path) { size_t size = 0; @@ -187,7 +192,7 @@ api::FuncGraphPtr TfliteModelParser::Parse(const converter::ConverterParameters return nullptr; } - auto func_graph = std::dynamic_pointer_cast(res_graph_); + auto func_graph = ConvertGraph(res_graph_); MS_CHECK_TRUE_RET(func_graph != nullptr, nullptr); if ((status = CommonAnfAdjust(func_graph)) != RET_OK) { MS_LOG(ERROR) << "AdjustForAnf failed."; @@ -250,7 +255,7 @@ STATUS TfliteModelParser::ConvertTfliteGraph() { // record the function graph if (idx == kMainGraphIndex) { - res_graph_ = func_graph; + res_graph_ = api::MakeShared(func_graph); } else { status = BuildSubFuncGraphMap(idx, func_graph, subgraph_name); if (status != RET_OK) { @@ -291,10 +296,10 @@ STATUS TfliteModelParser::ConvertOps(const std::unique_ptr &t op_idx++; // parse primitive MS_LOG(INFO) << "parse node :" << op_name; - ops::PrimitiveC *primitive_c = nullptr; + ops::PrimitiveCPtr primitive_c; auto node_parser = registry::NodeParserRegistry::GetNodeParser(kFmkTypeTflite, op_type); if (node_parser != nullptr) { - primitive_c = node_parser->Parse(op, tflite_subgraph, tflite_model_); + primitive_c = node_parser->Parse(op, tflite_subgraph, tflite_model_)->GetPrim(); } else { auto node_parser_builtin = TfliteNodeParserRegistry::GetInstance()->GetNodeParser(tflite_op_type); if (node_parser_builtin == nullptr) { @@ -311,7 +316,7 @@ STATUS TfliteModelParser::ConvertOps(const std::unique_ptr &t std::vector op_inputs; if (primitive_c != nullptr) { - auto value_node = NewValueNode(std::shared_ptr(primitive_c)); + auto value_node = NewValueNode(primitive_c); MSLITE_CHECK_PTR(value_node); op_inputs = {value_node}; } else { @@ -416,7 +421,7 @@ STATUS TfliteModelParser::SetTensorQuantParam(const std::unique_ptr &op, const std::unique_ptr &tflite_subgraph, - ops::PrimitiveC *primitive_c) { + PrimitiveCPtr primitive_c) { MS_ASSERT(tflite_subgraph != nullptr); if (op == nullptr) { MS_LOG(ERROR) << "tflite op is null, get quant params failed."; @@ -526,7 +531,9 @@ STATUS TfliteModelParser::ConvertGraphOutputs(const std::unique_ptrGetPrim(); + MSLITE_CHECK_PTR(make_tuple_prim_c); + auto make_tuple_prim = NewValueNode(make_tuple_prim_c); MSLITE_CHECK_PTR(make_tuple_prim); std::vector make_tuple_inputs = output_nodes; make_tuple_inputs.insert(make_tuple_inputs.begin(), make_tuple_prim); @@ -539,7 +546,9 @@ STATUS TfliteModelParser::ConvertGraphOutputs(const std::unique_ptrGetPrim(); + MSLITE_CHECK_PTR(return_prim_c); + auto value_node = NewValueNode(return_prim_c); MSLITE_CHECK_PTR(value_node); std::vector op_inputs{value_node}; op_inputs.emplace_back(make_tuple_cnode); @@ -556,7 +565,9 @@ STATUS TfliteModelParser::ConvertGraphOutputs(const std::unique_ptroutputs.front() < 0 ? static_cast(tflite_subgraph->outputs.front() + tflite_subgraph->tensors.size()) : static_cast(tflite_subgraph->outputs.front()); - auto value_node = NewValueNode(returnPrim); + auto return_prim_c = returnPrim->GetPrim(); + MSLITE_CHECK_PTR(return_prim_c); + auto value_node = NewValueNode(return_prim_c); MSLITE_CHECK_PTR(value_node); std::vector op_inputs{value_node}; MS_CHECK_TRUE_RET(anf_node_map->find(output_idx) != anf_node_map->end(), RET_NOT_FIND_OP); @@ -623,7 +634,7 @@ STATUS TfliteModelParser::ControlFlowNodePostProcess() { if (control_flow_map_.empty()) { return RET_OK; } - auto func_graph = std::dynamic_pointer_cast(res_graph_); + auto func_graph = ConvertGraph(res_graph_); MS_CHECK_TRUE_RET(func_graph != nullptr, RET_ERROR); static auto root_func_manager = Manage(func_graph); for (auto &node_vs_graph : control_flow_map_) { @@ -643,7 +654,7 @@ STATUS TfliteModelParser::ControlFlowNodePostProcess() { MSLITE_CHECK_PTR(second_value_node); auto inputs = control_flow_node->inputs(); inputs.insert(inputs.begin() + 1, {first_value_node, second_value_node}); - auto new_node = res_graph_->NewCNode(inputs); // must create new node, otherwise node_users won't update + auto new_node = func_graph->NewCNode(inputs); // must create new node, otherwise node_users won't update if (new_node == nullptr) { MS_LOG(ERROR) << "new node failed"; return RET_ERROR; @@ -824,7 +835,9 @@ STATUS TfliteModelParser::ConvertOutputTensor(const std::unique_ptrGetPrim(); + MSLITE_CHECK_PTR(tuple_get_item_prim_c); + auto tuple_get_item_prim = NewValueNode(tuple_get_item_prim_c); MSLITE_CHECK_PTR(tuple_get_item_prim); auto get_item_value = NewValueNode(MakeValue(op_idx)); MSLITE_CHECK_PTR(get_item_value); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.h index 417dd91407..7997e25c00 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_model_parser.h @@ -70,7 +70,7 @@ class TfliteModelParser : public converter::ModelParser { const CNodePtr &dst_cnode, std::unordered_map *anf_node_map); static STATUS ConvertOpQuantParams(const std::unique_ptr &op, const std::unique_ptr &tflite_subgraph, - ops::PrimitiveC *primitive_c); + PrimitiveCPtr primitive_c); static STATUS SetTensorQuantParam(const std::unique_ptr &tflite_tensor, std::vector *quant_params, int round_type = 1); STATUS TfliteOpVerify(const std::unique_ptr &subgraph, const size_t operator_codes_size, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_node_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_node_parser.h index ba83aecfe9..5c5db9b1ab 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_node_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_node_parser.h @@ -31,6 +31,9 @@ #include "tools/converter/parser/tflite/tflite_util.h" #include "ops/primitive_c.h" #include "src/common/log_util.h" +#include "tools/converter/parser/parser_utils.h" +#include "ops/op_utils.h" +#include "ops/op_name.h" namespace mindspore { namespace lite { @@ -40,9 +43,9 @@ class TfliteNodeParser { virtual ~TfliteNodeParser() = default; - virtual ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { + virtual PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { return nullptr; } diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.cc index 55b9eac2b4..803c98b167 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteOneHotParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteOneHotParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -35,7 +35,7 @@ ops::PrimitiveC *TfliteOneHotParser::Parse(const std::unique_ptrset_axis(tflite_attr->axis); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteOneHotParser(tflite::BuiltinOperator_ONE_HOT, new TfliteOneHotParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.h index 48147e1f8a..2c0ac51bc5 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_one_hot_parser.h @@ -31,9 +31,9 @@ class TfliteOneHotParser : public TfliteNodeParser { ~TfliteOneHotParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.cc index 74242447cd..1a8a609105 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.cc @@ -27,9 +27,9 @@ namespace { constexpr int kTFlitePadInputSize = 3; constexpr int kTFlitePaddingIndex = 1; } // namespace -ops::PrimitiveC *TflitePadParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TflitePadParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -89,7 +89,7 @@ ops::PrimitiveC *TflitePadParser::Parse(const std::unique_ptr return nullptr; } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tflitePadParser(tflite::BuiltinOperator_PAD, new TflitePadParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.h index 96afd9b8a5..3af6b9e7f1 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_pad_parser.h @@ -31,9 +31,9 @@ class TflitePadParser : public TfliteNodeParser { ~TflitePadParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.cc index 15575eb34d..6f595a6deb 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.cc @@ -24,9 +24,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteAvgPoolParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteAvgPoolParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_TRUE_RET(!tflite_op->inputs.empty(), nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -58,12 +58,12 @@ ops::PrimitiveC *TfliteAvgPoolParser::Parse(const std::unique_ptrset_pad(params); } - return prim.release(); + return prim->GetPrim(); } -ops::PrimitiveC *TfliteMaxPoolParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteMaxPoolParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_TRUE_RET(!tflite_op->inputs.empty(), nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -95,7 +95,7 @@ ops::PrimitiveC *TfliteMaxPoolParser::Parse(const std::unique_ptrset_pad(params); } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteMeanPoolingParser(tflite::BuiltinOperator_AVERAGE_POOL_2D, new TfliteAvgPoolParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.h index 69431ba32d..5193ce3521 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_pooling_parser.h @@ -31,9 +31,9 @@ class TfliteAvgPoolParser : public TfliteNodeParser { ~TfliteAvgPoolParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; class TfliteMaxPoolParser : public TfliteNodeParser { @@ -42,9 +42,9 @@ class TfliteMaxPoolParser : public TfliteNodeParser { ~TfliteMaxPoolParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.cc index 2d764bb6c0..ba51a5f549 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteQuantizeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteQuantizeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_TRUE_RET(!tflite_op->inputs.empty(), nullptr); MS_CHECK_TRUE_RET(!tflite_op->outputs.empty(), nullptr); const auto &in_tensor = tflite_subgraph->tensors[tflite_op->inputs[FIRST_INPUT]]; @@ -45,15 +45,17 @@ ops::PrimitiveC *TfliteQuantizeParser::Parse(const std::unique_ptrset_src_t(in_tensor_type); prim->set_dst_t(out_tensor_type); - return prim.release(); + return prim->GetPrim(); } else { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); auto dstT = GetTfliteDataType(out_tensor->type); auto dst_value = MakeValue(static_cast(dstT)); MS_CHECK_TRUE_RET(dst_value != nullptr, nullptr); - prim->AddAttr("to", dst_value); - return prim.release(); + prim_c->AddAttr("to", dst_value); + return prim->GetPrim(); } } diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.h index 1f73fe04b6..f588fd135a 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_quantize_parser.h @@ -30,9 +30,9 @@ class TfliteQuantizeParser : public TfliteNodeParser { ~TfliteQuantizeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.cc index b10d74c8ac..319b9f8d9c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteRangeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteRangeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize2, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -50,7 +50,7 @@ ops::PrimitiveC *TfliteRangeParser::Parse(const std::unique_ptrset_delta(delta.front()); } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteRangeParser(tflite::BuiltinOperator_RANGE, new TfliteRangeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.h index 854ccf8115..cef04878c5 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_range_parser.h @@ -31,9 +31,9 @@ class TfliteRangeParser : public TfliteNodeParser { ~TfliteRangeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.cc index ff5c9b0b04..0d29aa8dbb 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteRankParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteRankParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteRankParser(tflite::BuiltinOperator_RANK, new TfliteRankParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.h index a6c10961bd..60c3e068f8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_rank_parser.h @@ -31,9 +31,9 @@ class TfliteRankParser : public TfliteNodeParser { ~TfliteRankParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.cc index dbcdaadfc9..d05c50b26c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteReduceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteReduceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -53,7 +53,7 @@ ops::PrimitiveC *TfliteReduceParser::Parse(const std::unique_ptrGetPrim(); } TfliteNodeRegister g_TfliteSumParser(tflite::BuiltinOperator_SUM, new TfliteReduceParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.h index d2b1e1e8a6..822aa7a7a8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reduce_parser.h @@ -30,9 +30,9 @@ class TfliteReduceParser : public TfliteNodeParser { ~TfliteReduceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace mindspore::lite diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.cc index 762470136c..4b068d1992 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.cc @@ -22,12 +22,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteReshapeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tfliteModel) { +PrimitiveCPtr TfliteReshapeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tfliteModel) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, nullptr); std::vector shape; const auto &tflite_attr = tflite_op->builtin_options.AsReshapeOptions(); if (tflite_attr != nullptr) { @@ -37,10 +38,10 @@ ops::PrimitiveC *TfliteReshapeParser::Parse(const std::unique_ptrAddAttr("shape", value_ptr); + prim_c->AddAttr("shape", value_ptr); } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteReshapeParser(tflite::BuiltinOperator_RESHAPE, new TfliteReshapeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.h index ab04074f3a..84126d1594 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reshape_parser.h @@ -31,9 +31,9 @@ class TfliteReshapeParser : public TfliteNodeParser { ~TfliteReshapeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.cc index 7b346adecf..76e2eed7a3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.cc @@ -24,9 +24,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteResizeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteResizeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize1, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -67,7 +67,7 @@ ops::PrimitiveC *TfliteResizeParser::Parse(const std::unique_ptrhalf_pixel_centers) { prim->set_coordinate_transform_mode(mindspore::CoordinateTransformMode::HALF_PIXEL); prim->set_cubic_coeff(-0.5f); - prim->AddAttr("half_pixel_centers", MakeValue(true)); + prim->AddAttr("half_pixel_centers", api::MakeValue(true)); } prim->set_method(mindspore::ResizeMethod::NEAREST); prim->set_nearest_mode(mindspore::NearestMode::NORMAL); @@ -88,7 +88,7 @@ ops::PrimitiveC *TfliteResizeParser::Parse(const std::unique_ptrset_new_width(dims.at(1)); } - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteResizeBilinearParser(tflite::BuiltinOperator_RESIZE_BILINEAR, new TfliteResizeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.h index a29ea12cfc..c8e196d0f8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_resize_parser.h @@ -30,9 +30,9 @@ class TfliteResizeParser : public TfliteNodeParser { ~TfliteResizeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace mindspore::lite diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.cc index 65db4595b1..0ef4f24789 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteReverseParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteReverseParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize1, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -35,7 +35,7 @@ ops::PrimitiveC *TfliteReverseParser::Parse(const std::unique_ptrset_axis(axis); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteReverseParser(tflite::BuiltinOperator_REVERSE_V2, new TfliteReverseParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.h index 9929206d11..536d5a44b6 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_parser.h @@ -31,9 +31,9 @@ class TfliteReverseParser : public TfliteNodeParser { ~TfliteReverseParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.cc index 433b3cbd4c..8f7c29bbe6 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteReverseSequenceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteReverseSequenceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -38,7 +38,7 @@ ops::PrimitiveC *TfliteReverseSequenceParser::Parse(const std::unique_ptrset_seq_dim(tflite_attr->seq_dim); prim->set_batch_dim(tflite_attr->batch_dim); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteReverseSequenceParser(tflite::BuiltinOperator_REVERSE_SEQUENCE, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.h index be47c7adc2..99129b190d 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_reverse_sequence_parser.h @@ -31,9 +31,9 @@ class TfliteReverseSequenceParser : public TfliteNodeParser { ~TfliteReverseSequenceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.cc index 766c3da627..8ffd5128bc 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteScatterNdParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteScatterNdParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteScatterNdParser(tflite::BuiltinOperator_SCATTER_ND, new TfliteScatterNdParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.h index 440ded28a4..191e792095 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_scatter_nd_parser.h @@ -31,9 +31,9 @@ class TfliteScatterNdParser : public TfliteNodeParser { ~TfliteScatterNdParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.cc index 7e47f8dc0b..b31441a7c3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteShapeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteShapeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteShapeParser(tflite::BuiltinOperator_SHAPE, new TfliteShapeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.h index 982afe2601..b840d78ac0 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_shape_parser.h @@ -31,9 +31,9 @@ class TfliteShapeParser : public TfliteNodeParser { ~TfliteShapeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.cc index 88d41c1a65..c4fbd57d3e 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSkipGramParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSkipGramParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -37,7 +37,7 @@ ops::PrimitiveC *TfliteSkipGramParser::Parse(const std::unique_ptrset_max_skip_size(tflite_attr->max_skip_size); prim->set_ngram_size(tflite_attr->ngram_size); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSkiGramParser(tflite::BuiltinOperator_SKIP_GRAM, new TfliteSkipGramParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.h index c40b81bebb..734f610717 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_skip_gram_parser.h @@ -31,9 +31,9 @@ class TfliteSkipGramParser : public TfliteNodeParser { ~TfliteSkipGramParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.cc index 42cf92dacc..b56350f7b3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSliceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSliceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize1, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -41,7 +41,7 @@ ops::PrimitiveC *TfliteSliceParser::Parse(const std::unique_ptrset_axes(axes); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSliceParser(tflite::BuiltinOperator_SLICE, new TfliteSliceParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.h index b08b57e4a1..c9ca9a8df6 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_slice_parser.h @@ -31,9 +31,9 @@ class TfliteSliceParser : public TfliteNodeParser { ~TfliteSliceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.cc index 802314865d..c36abbe111 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.cc @@ -22,15 +22,15 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSoftmaxParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSoftmaxParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_axis({-1}); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSoftmaxParser(tflite::BuiltinOperator_SOFTMAX, new TfliteSoftmaxParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.h index bbf07373fe..24ea064c00 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_softmax_parser.h @@ -31,9 +31,9 @@ class TfliteSoftmaxParser : public TfliteNodeParser { ~TfliteSoftmaxParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.cc index 1f4eafd510..4c6fb02b95 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSpaceToBatchNDParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSpaceToBatchNDParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize2, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -44,7 +44,7 @@ ops::PrimitiveC *TfliteSpaceToBatchNDParser::Parse(const std::unique_ptrset_paddings(paddings); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSpaceToBatchNDParser(tflite::BuiltinOperator_SPACE_TO_BATCH_ND, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.h index f91aff3cd1..e1a10a90e2 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_batch_nd_parser.h @@ -31,9 +31,9 @@ class TfliteSpaceToBatchNDParser : public TfliteNodeParser { ~TfliteSpaceToBatchNDParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.cc index 5003ccdc5a..94fa6041f4 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSpaceToDepthParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSpaceToDepthParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -38,7 +38,7 @@ ops::PrimitiveC *TfliteSpaceToDepthParser::Parse(const std::unique_ptrset_block_size(tflite_attr->block_size); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSpaceToDepthParser(tflite::BuiltinOperator_SPACE_TO_DEPTH, new TfliteSpaceToDepthParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.h index c1d63cf7cd..075eb19e33 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_space_to_depth_parser.h @@ -31,9 +31,9 @@ class TfliteSpaceToDepthParser : public TfliteNodeParser { ~TfliteSpaceToDepthParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.cc index 0553617c7e..ab490c400e 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.cc @@ -23,12 +23,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSparseToDenseParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSparseToDenseParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSparseToDenseParser(tflite::BuiltinOperator_SPARSE_TO_DENSE, diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.h index 89bd7a0446..aa66ecdfd4 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_sparse_to_dense_parser.h @@ -31,9 +31,9 @@ class TfliteSparseToDenseParser : public TfliteNodeParser { ~TfliteSparseToDenseParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.cc index 188178e4f5..1c40141548 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSplitParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSplitParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize1, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -76,7 +76,7 @@ ops::PrimitiveC *TfliteSplitParser::Parse(const std::unique_ptrset_size_splits(size_splits); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSplitParser(tflite::BuiltinOperator_SPLIT, new TfliteSplitParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.h index 310e507659..2bd7bb4a5d 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_split_parser.h @@ -31,9 +31,9 @@ class TfliteSplitParser : public TfliteNodeParser { ~TfliteSplitParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.cc index 19ba047adc..692d30952d 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSplitVParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSplitVParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { MS_CHECK_GE(tflite_op->inputs.size(), kInputSize2, nullptr); auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -70,7 +70,7 @@ ops::PrimitiveC *TfliteSplitVParser::Parse(const std::unique_ptrset_axis(static_cast(axis)); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSplitVParser(tflite::BuiltinOperator_SPLIT_V, new TfliteSplitVParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.h index 5c730e786d..b21a82391e 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_split_v_parser.h @@ -31,9 +31,9 @@ class TfliteSplitVParser : public TfliteNodeParser { ~TfliteSplitVParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.cc index f7f3fb0ac5..3f7e6a1eb8 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteSqueezeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteSqueezeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -40,7 +40,7 @@ ops::PrimitiveC *TfliteSqueezeParser::Parse(const std::unique_ptr(value); }); prim->set_axis(dims_vector); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteSqueezeParser(tflite::BuiltinOperator_SQUEEZE, new TfliteSqueezeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.h index 026bc35fbd..ea8e09e4c3 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_squeeze_parser.h @@ -31,9 +31,9 @@ class TfliteSqueezeParser : public TfliteNodeParser { ~TfliteSqueezeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.cc index db1e1da169..2f7b6e03e7 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteStackParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteStackParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -35,7 +35,7 @@ ops::PrimitiveC *TfliteStackParser::Parse(const std::unique_ptrset_axis(tflite_attr->axis); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteStackParser(tflite::BuiltinOperator_PACK, new TfliteStackParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.h index 1235ebc0b0..70c8afe4bc 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_stack_parser.h @@ -31,9 +31,9 @@ class TfliteStackParser : public TfliteNodeParser { ~TfliteStackParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.cc index f2e783212c..8c80f53674 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.cc @@ -22,9 +22,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteStridedSliceParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteStridedSliceParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -39,7 +39,7 @@ ops::PrimitiveC *TfliteStridedSliceParser::Parse(const std::unique_ptrset_new_axis_mask(tflite_attr->new_axis_mask); prim->set_shrink_axis_mask(tflite_attr->shrink_axis_mask); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteStridedSliceParser(tflite::BuiltinOperator_STRIDED_SLICE, new TfliteStridedSliceParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.h index be4d91c6a4..1fb485b04b 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_strided_slice_parser.h @@ -31,9 +31,9 @@ class TfliteStridedSliceParser : public TfliteNodeParser { ~TfliteStridedSliceParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.cc index ffd25ef748..607fa4270a 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.cc @@ -23,12 +23,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteTileParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteTileParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteTileParser(tflite::BuiltinOperator_TILE, new TfliteTileParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.h index d8ea251cb9..be4c43f4d9 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_tile_parser.h @@ -31,9 +31,9 @@ class TfliteTileParser : public TfliteNodeParser { ~TfliteTileParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.cc index 16d3f1d09b..b8c78e799c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.cc @@ -23,15 +23,15 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteTopKV2Parser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteTopKV2Parser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); prim->set_sorted(true); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteTopKV2Parser(tflite::BuiltinOperator_TOPK_V2, new TfliteTopKV2Parser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.h index 18631fbf2b..b9761af1d5 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_topk_v2_parser.h @@ -31,9 +31,9 @@ class TfliteTopKV2Parser : public TfliteNodeParser { ~TfliteTopKV2Parser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.cc index 175e2495d0..033973ccfb 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.cc @@ -22,12 +22,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteTransposeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteTransposeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteTransposeParser(tflite::BuiltinOperator_TRANSPOSE, new TfliteTransposeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.h index cadfba9ea9..748f4b4960 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_transpose_parser.h @@ -31,9 +31,9 @@ class TfliteTransposeParser : public TfliteNodeParser { ~TfliteTransposeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.cc index 9d3c0430a3..3888470859 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.cc @@ -23,12 +23,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteUniqueParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteUniqueParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteUniqueParser(tflite::BuiltinOperator_UNIQUE, new TfliteUniqueParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.h index 2314b67fa7..5eb7f38b09 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_unique_parser.h @@ -31,9 +31,9 @@ class TfliteUniqueParser : public TfliteNodeParser { ~TfliteUniqueParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.cc index 51e78810d6..7627af7116 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.cc @@ -23,9 +23,9 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteUnstackParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteUnstackParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -36,7 +36,7 @@ ops::PrimitiveC *TfliteUnstackParser::Parse(const std::unique_ptrset_axis(tflite_attr->axis); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteUnstackParser(tflite::BuiltinOperator_UNPACK, new TfliteUnstackParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.h index 545bea0745..b6af808760 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_unstack_parser.h @@ -31,9 +31,9 @@ class TfliteUnstackParser : public TfliteNodeParser { ~TfliteUnstackParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.cc index 9960c08c6d..5ce1b6ac50 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.cc @@ -23,13 +23,13 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteWhereParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteWhereParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteWhereParser(tflite::BuiltinOperator_WHERE, new TfliteWhereParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.h index f8c4996f85..9ae8ce0e29 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_where_parser.h @@ -31,9 +31,9 @@ class TfliteWhereParser : public TfliteNodeParser { ~TfliteWhereParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.cc index ff6a5c0164..d04d5b802c 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.cc @@ -23,10 +23,10 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteWhileParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { - auto prim = std::make_unique(); +PrimitiveCPtr TfliteWhileParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { + auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); const auto &tflite_attr = tflite_op->builtin_options.AsWhileOptions(); @@ -37,7 +37,7 @@ ops::PrimitiveC *TfliteWhileParser::Parse(const std::unique_ptrset_cond_subgraph_index(tflite_attr->cond_subgraph_index); prim->set_body_subgraph_index(tflite_attr->body_subgraph_index); - return prim.release(); + return prim; } TfliteNodeRegister g_tfliteWhileParser(tflite::BuiltinOperator_WHILE, new TfliteWhileParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.h index cea077ee45..bc1ee4f891 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_while_parser.h @@ -31,9 +31,9 @@ class TfliteWhileParser : public TfliteNodeParser { ~TfliteWhileParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.cc b/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.cc index df1462a8a3..896709c909 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.cc +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.cc @@ -23,12 +23,12 @@ namespace mindspore { namespace lite { -ops::PrimitiveC *TfliteZerosLikeParser::Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) { +PrimitiveCPtr TfliteZerosLikeParser::Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) { auto prim = std::make_unique(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); - return prim.release(); + return prim->GetPrim(); } TfliteNodeRegister g_tfliteZerosLikeParser(tflite::BuiltinOperator_ZEROS_LIKE, new TfliteZerosLikeParser()); diff --git a/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.h b/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.h index 201dcf7a4e..2ee1039a9f 100644 --- a/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.h +++ b/mindspore/lite/tools/converter/parser/tflite/tflite_zeros_like_parser.h @@ -31,9 +31,9 @@ class TfliteZerosLikeParser : public TfliteNodeParser { ~TfliteZerosLikeParser() override = default; - ops::PrimitiveC *Parse(const std::unique_ptr &tflite_op, - const std::unique_ptr &tflite_subgraph, - const std::unique_ptr &tflite_model) override; + PrimitiveCPtr Parse(const std::unique_ptr &tflite_op, + const std::unique_ptr &tflite_subgraph, + const std::unique_ptr &tflite_model) override; }; } // namespace lite } // namespace mindspore diff --git a/mindspore/lite/tools/converter/parser/unify_format.cc b/mindspore/lite/tools/converter/parser/unify_format.cc index 677dc06db8..aa9d60426b 100644 --- a/mindspore/lite/tools/converter/parser/unify_format.cc +++ b/mindspore/lite/tools/converter/parser/unify_format.cc @@ -14,12 +14,14 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/parser/unify_format.h" #include #include #include #include "tools/common/tensor_util.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace lite { diff --git a/mindspore/lite/tools/converter/parser/unused_node_remove_pass.cc b/mindspore/lite/tools/converter/parser/unused_node_remove_pass.cc index ded162da50..b1d3c4af05 100644 --- a/mindspore/lite/tools/converter/parser/unused_node_remove_pass.cc +++ b/mindspore/lite/tools/converter/parser/unused_node_remove_pass.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/parser/unused_node_remove_pass.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/bias_correction_strategy.cc b/mindspore/lite/tools/converter/quantizer/bias_correction_strategy.cc index 3af1ed1392..d93da4317c 100644 --- a/mindspore/lite/tools/converter/quantizer/bias_correction_strategy.cc +++ b/mindspore/lite/tools/converter/quantizer/bias_correction_strategy.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/bias_correction_strategy.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/calibrator.cc b/mindspore/lite/tools/converter/quantizer/calibrator.cc index c1248e8acc..c74ba18007 100644 --- a/mindspore/lite/tools/converter/quantizer/calibrator.cc +++ b/mindspore/lite/tools/converter/quantizer/calibrator.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/calibrator.h" #include #include "tools/converter/preprocess/image_preprocess.h" @@ -25,7 +26,7 @@ namespace mindspore::lite::quant { namespace { constexpr int kDefaultBinNumber = 2048; -} +} // namespace int Calibrator::RecordMaxMinValue(const std::vector &data, const std::unique_ptr &diverg_info) { auto ret = diverg_info->RecordMaxMinValueArray(data); diff --git a/mindspore/lite/tools/converter/quantizer/data_distribution.cc b/mindspore/lite/tools/converter/quantizer/data_distribution.cc index 41b81fe4b9..a23790325d 100644 --- a/mindspore/lite/tools/converter/quantizer/data_distribution.cc +++ b/mindspore/lite/tools/converter/quantizer/data_distribution.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/data_distribution.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/debug_info_manager.cc b/mindspore/lite/tools/converter/quantizer/debug_info_manager.cc index 80dbd143d0..e2bacf9ed1 100644 --- a/mindspore/lite/tools/converter/quantizer/debug_info_manager.cc +++ b/mindspore/lite/tools/converter/quantizer/debug_info_manager.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include #include #include "tools/converter/quantizer/debug_info_manager.h" @@ -27,7 +28,7 @@ namespace mindspore::lite { namespace { constexpr int kNumUsPerMs = 1000; -} +} // namespace std::string DebugInfoManager::ParseInOutTensorToString(InOutFlag in_out_flag) { switch (in_out_flag) { case INPUT: diff --git a/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc b/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc index 863abb9a9d..ef2211754b 100644 --- a/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc +++ b/mindspore/lite/tools/converter/quantizer/dynamic_quantizer.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/dynamic_quantizer.h" #include "tools/converter/quantizer/weight_quantizer.h" #include "tools/converter/quantizer/insert_quant_node_manager.h" diff --git a/mindspore/lite/tools/converter/quantizer/full_quant_quantizer.cc b/mindspore/lite/tools/converter/quantizer/full_quant_quantizer.cc index a06e95f0a6..f2c5b095b9 100644 --- a/mindspore/lite/tools/converter/quantizer/full_quant_quantizer.cc +++ b/mindspore/lite/tools/converter/quantizer/full_quant_quantizer.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/full_quant_quantizer.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/huffman_encode.cc b/mindspore/lite/tools/converter/quantizer/huffman_encode.cc index 424a2a968f..7f183f6a04 100644 --- a/mindspore/lite/tools/converter/quantizer/huffman_encode.cc +++ b/mindspore/lite/tools/converter/quantizer/huffman_encode.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/huffman_encode.h" #include "src/weight_decoder.h" #include "tools/converter/quantizer/quantize_util.h" diff --git a/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc b/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc index 2f71095244..07259ce3b0 100644 --- a/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc +++ b/mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "mindspore/lite/tools/converter/quantizer/insert_quant_node_manager.h" #include #include @@ -43,8 +44,10 @@ ValueNodePtr InsertQuantNodeManager::NewQuantCastValueNode(int src_type, int dst quant_params_holder->set_input_quant_param(i, quant_params_in); quant_params_holder->set_output_quant_param(i, quant_params_in); } - prim_c->AddAttr("quant_params", quant_params_holder); - return NewValueNode(prim_c); + auto prim = prim_c->GetPrim(); + MS_CHECK_TRUE_MSG(prim != nullptr, nullptr, "prim is nullptr"); + prim->AddAttr("quant_params", quant_params_holder); + return NewValueNode(prim); } int InsertQuantNodeManager::InsertCastNode(const FuncGraphPtr &graph, const CNodePtr &cnode, size_t input_index, @@ -164,9 +167,10 @@ int InsertQuantNodeManager::InsertQuantDtypeCastNode(const FuncGraphPtr &graph) int InsertQuantNodeManager::InsertDynamicQuantWithIndex(const FuncGraphPtr &graph, const CNodePtr &cnode, size_t index) { - auto primitive_c = std::make_shared(); - primitive_c->set_dst_type(dst_type_); - primitive_c->set_symmetric(symmetric_); + auto primitive = std::make_shared(); + auto primitive_c = primitive->GetPrim(); + primitive->set_dst_type(dst_type_); + primitive->set_symmetric(symmetric_); auto dynamic_quant_cnode = graph->NewCNode(primitive_c, {cnode->input(index)}); auto name = cnode->fullname_with_scope() + "_dynamic_cast_node_" + to_string(index); dynamic_quant_cnode->set_fullname_with_scope(name); diff --git a/mindspore/lite/tools/converter/quantizer/parameter_tunner.cc b/mindspore/lite/tools/converter/quantizer/parameter_tunner.cc index 91585cfe07..ff0308df12 100644 --- a/mindspore/lite/tools/converter/quantizer/parameter_tunner.cc +++ b/mindspore/lite/tools/converter/quantizer/parameter_tunner.cc @@ -13,7 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ - +#define USE_DEPRECATED_API #include "tools/converter/quantizer/parameter_tunner.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/quant_helper/attention_quant_type_determiner.cc b/mindspore/lite/tools/converter/quantizer/quant_helper/attention_quant_type_determiner.cc index f7a5cf2a9d..201c28bede 100644 --- a/mindspore/lite/tools/converter/quantizer/quant_helper/attention_quant_type_determiner.cc +++ b/mindspore/lite/tools/converter/quantizer/quant_helper/attention_quant_type_determiner.cc @@ -13,6 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quant_helper/attention_quant_type_determiner.h" #include "tools/converter/quantizer/quant_helper/conv_quant_param_propogator.h" #include "tools/converter/quantizer/quantize_util.h" diff --git a/mindspore/lite/tools/converter/quantizer/quant_helper/bias_add_quant_param_propogator.cc b/mindspore/lite/tools/converter/quantizer/quant_helper/bias_add_quant_param_propogator.cc index af70c13c2c..84eab8837d 100644 --- a/mindspore/lite/tools/converter/quantizer/quant_helper/bias_add_quant_param_propogator.cc +++ b/mindspore/lite/tools/converter/quantizer/quant_helper/bias_add_quant_param_propogator.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quant_helper/bias_add_quant_param_propogator.h" #include "mindspore/core/ir/dtype/type_id.h" #include "mindspore/core/utils/log_adapter.h" diff --git a/mindspore/lite/tools/converter/quantizer/quant_helper/concat_quant_param_propogator.cc b/mindspore/lite/tools/converter/quantizer/quant_helper/concat_quant_param_propogator.cc index d3b537d9c1..59ec08bcf5 100644 --- a/mindspore/lite/tools/converter/quantizer/quant_helper/concat_quant_param_propogator.cc +++ b/mindspore/lite/tools/converter/quantizer/quant_helper/concat_quant_param_propogator.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quant_helper/concat_quant_param_propogator.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/quant_helper/conv_quant_type_determiner.cc b/mindspore/lite/tools/converter/quantizer/quant_helper/conv_quant_type_determiner.cc index 9b5dff3994..1f86945389 100644 --- a/mindspore/lite/tools/converter/quantizer/quant_helper/conv_quant_type_determiner.cc +++ b/mindspore/lite/tools/converter/quantizer/quant_helper/conv_quant_type_determiner.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quant_helper/conv_quant_type_determiner.h" #include "tools/converter/quantizer/quantize_util.h" #include "mindspore/core/utils/log_adapter.h" diff --git a/mindspore/lite/tools/converter/quantizer/quant_helper/matmul_quant_type_determiner.cc b/mindspore/lite/tools/converter/quantizer/quant_helper/matmul_quant_type_determiner.cc index aa2d051e08..bfce2dd449 100644 --- a/mindspore/lite/tools/converter/quantizer/quant_helper/matmul_quant_type_determiner.cc +++ b/mindspore/lite/tools/converter/quantizer/quant_helper/matmul_quant_type_determiner.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quant_helper/matmul_quant_type_determiner.h" #include "tools/converter/quantizer/quantize_util.h" #include "mindspore/core/utils/log_adapter.h" diff --git a/mindspore/lite/tools/converter/quantizer/quant_strategy.cc b/mindspore/lite/tools/converter/quantizer/quant_strategy.cc index 8ee3038888..b787d2486f 100644 --- a/mindspore/lite/tools/converter/quantizer/quant_strategy.cc +++ b/mindspore/lite/tools/converter/quantizer/quant_strategy.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quant_strategy.h" #include #include "tools/converter/quantizer/quantize_util.h" @@ -21,6 +22,7 @@ #include "src/common/log_adapter.h" #include "src/common/log_util.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::lite::quant { bool QuantStrategy::CanTensorQuantized(const CNodePtr &cnode, const AnfNodePtr &input_node, int preferred_dim) { diff --git a/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc b/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc index d16cae955a..f455615cef 100644 --- a/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc +++ b/mindspore/lite/tools/converter/quantizer/quantization_optimizer.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/quantization_optimizer.h" #include #include diff --git a/mindspore/lite/tools/converter/quantizer/quantize_util.cc b/mindspore/lite/tools/converter/quantizer/quantize_util.cc index 1756fc2c66..48db941efa 100644 --- a/mindspore/lite/tools/converter/quantizer/quantize_util.cc +++ b/mindspore/lite/tools/converter/quantizer/quantize_util.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "mindspore/lite/tools/converter/quantizer/quantize_util.h" #include #include @@ -34,6 +35,8 @@ #include "securec/include/securec.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "tools/optimizer/common/format_utils.h" +#include "ops/op_utils.h" using std::string; using std::vector; @@ -382,7 +385,7 @@ int UpdateTensorDataAndSize(const AnfNodePtr &node, const tensor::TensorPtr &wei int GetMatMulPreferredDim(const PrimitivePtr &primitive, int input_index, const std::vector &dims) { size_t last_first_index = dims.size() - 1; size_t last_second_index = dims.size() - 2; - auto matmul_prim = primitive->cast>(); + auto matmul_prim = api::MakeShared(primitive); MS_ASSERT(matmul_prim != nullptr); // For MatMul A if (input_index == 0) { @@ -404,7 +407,7 @@ int GetMatMulPreferredDim(const PrimitivePtr &primitive, int input_index, const } int GetDeConvPreferredDim(const PrimitivePtr &primitive, const std::vector &dims) { - auto prim = primitive->cast>(); + auto prim = api::MakeShared(primitive); MS_ASSERT(prim != nullptr); // For MatMul A if (prim->get_in_channel() == prim->get_group() && prim->get_out_channel() == prim->get_group()) { diff --git a/mindspore/lite/tools/converter/quantizer/quantizer.h b/mindspore/lite/tools/converter/quantizer/quantizer.h index 7ce1e599b9..5f516aa031 100644 --- a/mindspore/lite/tools/converter/quantizer/quantizer.h +++ b/mindspore/lite/tools/converter/quantizer/quantizer.h @@ -16,7 +16,6 @@ #ifndef MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_QUANTIZER_H #define MINDSPORE_LITE_TOOLS_CONVERTER_QUANTIZER_QUANTIZER_H - #include #include #include diff --git a/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc b/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc index c4808b17f3..6c737865ab 100644 --- a/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc +++ b/mindspore/lite/tools/converter/quantizer/weight_quantizer.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/converter/quantizer/weight_quantizer.h" #include #include @@ -21,6 +22,7 @@ #include #include "tools/optimizer/common/gllo_utils.h" #include "src/common/log_util.h" +#include "api/ir/func_graph.h" namespace mindspore::lite::quant { WeightQuantizer::~WeightQuantizer() { @@ -89,7 +91,7 @@ int WeightQuantizer::DoCNodeWeightQuant(const FuncGraphPtr &func_graph, const CN CHECK_NULL_RETURN(cnode); auto primitive = GetValueNode(cnode->input(0)); CHECK_NULL_RETURN(primitive); - auto manager = api::FuncGraphManager::Manage(func_graph, true); + auto manager = deprecated::api::FuncGraphManager::Manage(func_graph, true); CHECK_NULL_RETURN(manager); for (auto idx : weight_indices) { auto input = cnode->input(idx); diff --git a/mindspore/lite/tools/optimizer/common/gllo_utils.cc b/mindspore/lite/tools/optimizer/common/gllo_utils.cc index e2e7b67914..fc46af6ce1 100644 --- a/mindspore/lite/tools/optimizer/common/gllo_utils.cc +++ b/mindspore/lite/tools/optimizer/common/gllo_utils.cc @@ -33,6 +33,7 @@ #include "nnacl/op_base.h" #include "src/common/log_util.h" #include "tools/converter/parser/parser_utils.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -909,7 +910,8 @@ CNodePtr GenTransposeNode(const FuncGraphPtr &func_graph, const AnfNodePtr &inpu MS_CHECK_TRUE_RET(input_node != nullptr, nullptr); auto perm_node = BuildIntVecParameterNode(func_graph, perm, cnode_name + "_perm"); MS_ASSERT(perm_node != nullptr); - auto trans_prim = std::make_shared(); + ops::Transpose transpose_node; + auto trans_prim = transpose_node.GetPrim(); MS_CHECK_TRUE_RET(trans_prim != nullptr, nullptr); auto cnode = func_graph->NewCNode(trans_prim, {input_node, perm_node}); MS_ASSERT(cnode != nullptr); @@ -940,7 +942,8 @@ CNodePtr GenGatherNode(const FuncGraphPtr &func_graph, const AnfNodePtr &input_n MS_LOG(ERROR) << "make indices node failed."; return nullptr; } - auto gather_prim = std::make_shared(); + ops::Gather gather_node; + auto gather_prim = gather_node.GetPrim(); MS_CHECK_TRUE_RET(gather_prim != nullptr, nullptr); auto cnode = func_graph->NewCNode(gather_prim, {input_node, indices_node, axis_node}); MS_ASSERT(cnode != nullptr); @@ -965,7 +968,9 @@ CNodePtr GenTupleGetItemNode(const FuncGraphPtr &func_graph, const CNodePtr &inp MS_CHECK_TRUE_RET(tuple_get_item_prim != nullptr, nullptr); auto second_input = NewValueNode(MakeValue(index)); MS_CHECK_TRUE_RET(second_input != nullptr, nullptr); - auto tuple_cnode = func_graph->NewCNode(tuple_get_item_prim, {input, second_input}); + auto tuple_get_item_prim_c = tuple_get_item_prim->GetPrim(); + MS_CHECK_TRUE_RET(tuple_get_item_prim_c != nullptr, nullptr); + auto tuple_cnode = func_graph->NewCNode(tuple_get_item_prim_c, {input, second_input}); MS_ASSERT(tuple_cnode != nullptr); tuple_cnode->set_fullname_with_scope(input->fullname_with_scope() + "_getitem_" + std::to_string(index)); return tuple_cnode; @@ -991,15 +996,15 @@ STATUS FetchShapeFromAbstract(const abstract::AbstractBasePtr &abstract, ShapeVe bool IsTrainOp(const CNodePtr &cnode) { auto prim = GetValueNode(cnode->input(0)); - auto cnode_type = prim->type_name(); + auto cnode_type = prim->name(); // optimizer op if (cnode_type == "Adam" || cnode_type == "SGD" || cnode_type == "ApplyMomentum") { return true; } // loss op - if (cnode_type == "SoftmaxCrossEntropyWithLogits" || cnode_type == "SpareSoftmaxCrossEntropyWithLogits" || + if (cnode_type == "SoftmaxCrossEntropyWithLogits" || cnode_type == "SparseSoftmaxCrossEntropyWithLogits" || cnode_type == "SmoothL1Loss" || cnode_type == "SmoothL1LossGrad" || - cnode_type == "SigmoidCrossEntropyWithLogits" || cnode_type == "SigmoidCrossEntropyWithLogpitsGrad") { + cnode_type == "SigmoidCrossEntropyWithLogits" || cnode_type == "SigmoidCrossEntropyWithLogitsGrad") { return true; } // grad op diff --git a/mindspore/lite/tools/optimizer/common/gllo_utils.h b/mindspore/lite/tools/optimizer/common/gllo_utils.h index 707fcc06d1..50487f13e5 100644 --- a/mindspore/lite/tools/optimizer/common/gllo_utils.h +++ b/mindspore/lite/tools/optimizer/common/gllo_utils.h @@ -16,7 +16,9 @@ #ifndef MINDSPORE_LITE_TOOLS_OPTIMIZER_COMMON_GLLO_UTILS_H_ #define MINDSPORE_LITE_TOOLS_OPTIMIZER_COMMON_GLLO_UTILS_H_ - +#ifndef USE_DEPRECATED_API +#define USE_DEPRECATED_API +#endif #include #include #include @@ -54,8 +56,7 @@ inline const std::vector kNH2NC = {0, 3, 1, 2}; inline const std::vector kNC2NH = {0, 2, 3, 1}; inline const PrimitivePtr kPrimMakeTupleV2 = std::make_shared("make_tuple"); inline const PrimitivePtr kPrimIdentity = std::make_shared("Identity"); -inline const PrimitivePtr kPrimConv2DBackpropInputFusion = - std::make_shared(ops::kNameConv2DBackpropInputFusion); +inline const PrimitivePtr kPrimConv2DBackpropInputFusion = std::make_shared("Conv2DBackpropInputFusion"); std::vector CastToInt(const ValuePtr &value); diff --git a/mindspore/lite/tools/optimizer/common/helper.cc b/mindspore/lite/tools/optimizer/common/helper.cc index 7b62298ea3..7bc32fd9af 100644 --- a/mindspore/lite/tools/optimizer/common/helper.cc +++ b/mindspore/lite/tools/optimizer/common/helper.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "backend/common/optimizer/helper.h" #include #include diff --git a/mindspore/lite/tools/optimizer/common/node_pass_extends.cc b/mindspore/lite/tools/optimizer/common/node_pass_extends.cc index 670ce500b3..a23e6aa1d4 100644 --- a/mindspore/lite/tools/optimizer/common/node_pass_extends.cc +++ b/mindspore/lite/tools/optimizer/common/node_pass_extends.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "backend/common/optimizer/node_pass.h" #include #include diff --git a/mindspore/lite/tools/optimizer/const_fold/fold_along_infershape.cc b/mindspore/lite/tools/optimizer/const_fold/fold_along_infershape.cc index 1adb4aac30..0fa4dd8dfd 100644 --- a/mindspore/lite/tools/optimizer/const_fold/fold_along_infershape.cc +++ b/mindspore/lite/tools/optimizer/const_fold/fold_along_infershape.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/const_fold/fold_along_infershape.h" #include #include "nnacl/op_base.h" diff --git a/mindspore/lite/tools/optimizer/const_fold/fold_utils.cc b/mindspore/lite/tools/optimizer/const_fold/fold_utils.cc index 6cbf45e38b..4037cc8bec 100644 --- a/mindspore/lite/tools/optimizer/const_fold/fold_utils.cc +++ b/mindspore/lite/tools/optimizer/const_fold/fold_utils.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/const_fold/fold_utils.h" #include #include diff --git a/mindspore/lite/tools/optimizer/const_fold/fold_with_infershape.cc b/mindspore/lite/tools/optimizer/const_fold/fold_with_infershape.cc index 50dbbe2cba..ec1f1e93f5 100644 --- a/mindspore/lite/tools/optimizer/const_fold/fold_with_infershape.cc +++ b/mindspore/lite/tools/optimizer/const_fold/fold_with_infershape.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/const_fold/fold_with_infershape.h" #include #include diff --git a/mindspore/lite/tools/optimizer/fisson/eliminate_concat_split.cc b/mindspore/lite/tools/optimizer/fisson/eliminate_concat_split.cc index 631cc2826d..2d1ee3c7a8 100644 --- a/mindspore/lite/tools/optimizer/fisson/eliminate_concat_split.cc +++ b/mindspore/lite/tools/optimizer/fisson/eliminate_concat_split.cc @@ -18,6 +18,7 @@ #include #include #include +#include "tools/common/node_util.h" #include "tools/optimizer/fisson/eliminate_concat_split.h" #include "schema/inner/model_generated.h" #include "include/common/utils/utils.h" @@ -28,6 +29,7 @@ #include "tools/optimizer/parallel/spliter.h" #include "src/common/log_util.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -72,9 +74,9 @@ int ConcatSplitEliminate(const FuncGraphPtr &func_graph, const CNodePtr &cnode) size_t pre_inputs_size = pre_cnode->inputs().size(); auto pre_inputs_node_size = static_cast(pre_inputs_size - 1); - auto pre_prim = GetValueNode>(pre_cnode->input(kAnfPrimitiveIndex)); + auto pre_prim = ops::GetOperator(pre_cnode->input(kAnfPrimitiveIndex)); MS_CHECK_TRUE_MSG(pre_prim != nullptr, lite::RET_ERROR, "pre_cnode is not a ops::Concat"); - auto prim = GetValueNode>(cnode->input(kAnfPrimitiveIndex)); + auto prim = ops::GetOperator(cnode->input(kAnfPrimitiveIndex)); MS_CHECK_TRUE_MSG(prim != nullptr, lite::RET_ERROR, "cnode is not a ops::SplitWithOverlap"); if (prim->get_number_split() != pre_inputs_node_size) { return RET_OK; @@ -138,7 +140,8 @@ int ConcatSplitEliminate(const FuncGraphPtr &func_graph, const CNodePtr &cnode) const BaseRef EliminateConcatSplit::DefinePattern() const { auto concat_var = std::make_shared(IsSpecifiedNode<&prim::kPrimConcat>); CHECK_NULL_RETURN(concat_var); - auto split_prim = std::make_shared(); + ops::SplitWithOverlap split_node; + auto split_prim = split_node.GetPrim(); CHECK_NULL_RETURN(split_prim); return VectorRef({split_prim, concat_var}); } diff --git a/mindspore/lite/tools/optimizer/fisson/fisson_util.cc b/mindspore/lite/tools/optimizer/fisson/fisson_util.cc index 0ff02c8da0..fa08ec2dc0 100644 --- a/mindspore/lite/tools/optimizer/fisson/fisson_util.cc +++ b/mindspore/lite/tools/optimizer/fisson/fisson_util.cc @@ -25,11 +25,12 @@ #include "tools/optimizer/parallel/split_strategy.h" #include "nnacl/op_base.h" #include "src/common/log_util.h" +#include "ops/op_utils.h" using mindspore::converter::FmkType; namespace mindspore { namespace opt { -std::vector GetSplitPadList(const std::shared_ptr &ori_conv_prim, int64_t input_h, +std::vector GetSplitPadList(const api::SharedPtr &ori_conv_prim, int64_t input_h, int64_t input_w) { if (ori_conv_prim == nullptr) { MS_LOG(DEBUG) << "input Conv2DFusion is nullptr"; @@ -114,7 +115,7 @@ bool CalSplitOutputShape(int64_t splited_axis_value, const SplitInfo *split_info } bool CalSplitInShape(const std::vector> &node_in_out_shapes, const SplitInfo *split_info, - const std::shared_ptr &ori_conv_prim, size_t index_node, + const api::SharedPtr &ori_conv_prim, size_t index_node, std::vector> *split_axis_inputs_shape, std::vector> *split_axis_reduce_inputs_shape) { MS_ASSERT(split_info != nullptr && ori_conv_prim != nullptr && split_axis_inputs_shape != nullptr && @@ -180,10 +181,12 @@ bool IsConv2D(const AnfNodePtr &node) { return (CheckPrimitiveType(node, prim::kPrimConv2D) || CheckPrimitiveType(node, prim::kPrimConv2DFusion)); } -std::shared_ptr CopyConvPrim(const std::shared_ptr &ori_conv_prim) { +api::SharedPtr CopyConvPrim(const api::SharedPtr &ori_conv_prim) { MS_CHECK_TRUE_MSG(ori_conv_prim != nullptr, nullptr, "input Conv2DFusion is nullptr"); - auto new_prim = std::make_shared(); + auto new_prim = api::MakeShared(); MS_CHECK_TRUE_MSG(new_prim != nullptr, nullptr, "create Conv2DFusion return nullptr"); + auto new_prim_c = new_prim->GetPrim(); + MS_CHECK_TRUE_MSG(new_prim_c != nullptr, nullptr, "create primic return nullptr"); new_prim->set_pad(ori_conv_prim->get_pad()); new_prim->set_in_channel(ori_conv_prim->get_in_channel()); new_prim->set_out_channel(ori_conv_prim->get_out_channel()); @@ -203,7 +206,7 @@ std::shared_ptr CopyConvPrim(const std::shared_ptrGetAttr(ops::kIsDepthWise); if (is_depth_value != nullptr) { bool is_depth_wise = GetValue(is_depth_value); - new_prim->AddAttr(ops::kIsDepthWise, MakeValue(is_depth_wise)); + new_prim_c->AddAttr(ops::kIsDepthWise, MakeValue(is_depth_wise)); } return new_prim; } @@ -260,7 +263,7 @@ bool UpdateSplitInfo(const FuncGraphPtr &func_graph, const std::vectorcast(); MS_ASSERT(conv_cnode != nullptr); - auto ori_conv_prim = GetValueNode>(conv_cnode->input(kAnfPrimitiveIndex)); + auto ori_conv_prim = ops::GetOperator(conv_cnode->input(kAnfPrimitiveIndex)); MS_CHECK_TRUE_RET(ori_conv_prim != nullptr, false); if (!CalSplitInShape(node_in_out_shapes, split_info, ori_conv_prim, index_node, &split_axis_inputs_shape, &split_axis_reduce_inputs_shape)) { @@ -332,10 +335,12 @@ AnfNodePtr CreateOutputsOfConcat(const FuncGraphPtr &func_graph, const AnfNodePt auto concat_prim = std::make_shared(); MS_CHECK_TRUE_MSG(concat_prim != nullptr, nullptr, "create ops::Concat return nullptr"); + auto concat_prim_c = concat_prim->GetPrim(); + MS_CHECK_TRUE_MSG(concat_prim_c != nullptr, nullptr, "create ops::concat_prim_c return nullptr"); concat_prim->set_axis(split_info.axis); // the inputs of concate are from the outputs of conv - auto concate_primitive = NewValueNode(concat_prim); + auto concate_primitive = NewValueNode(concat_prim_c); MS_CHECK_TRUE_MSG(concate_primitive != nullptr, nullptr, "create concate_primitive return nullptr"); std::vector concate_inputs = {concate_primitive}; for (size_t i = 0; i < static_cast(nodes_num); i++) { @@ -364,6 +369,8 @@ bool CreateOutputsOfSplitWithOverlap(const FuncGraphPtr &func_graph, const AnfNo // attr of split auto split_prim = std::make_shared(); MS_CHECK_TRUE_MSG(split_prim != nullptr, false, "create ops::SplitWithOverlap return nullptr"); + auto split_prim_c = split_prim->GetPrim(); + MS_CHECK_TRUE_MSG(split_prim != nullptr, false, "create ops::split_prim_c return nullptr"); split_prim->set_split_dim(split_info.axis); split_prim->set_number_split(split_info.out_num); split_prim->set_ratio(split_info.size_splits); @@ -372,7 +379,7 @@ bool CreateOutputsOfSplitWithOverlap(const FuncGraphPtr &func_graph, const AnfNo auto conv_cnode = conv_node->cast(); // the inputs of split is from the inputs of conv - auto split_primitive = NewValueNode(split_prim); + auto split_primitive = NewValueNode(split_prim_c); MS_CHECK_TRUE_MSG(split_primitive != nullptr, false, "create split_primitive return nullptr"); std::vector split_inputs = {split_primitive}; diff --git a/mindspore/lite/tools/optimizer/fisson/fisson_util.h b/mindspore/lite/tools/optimizer/fisson/fisson_util.h index 73c1aa9598..6984773b94 100644 --- a/mindspore/lite/tools/optimizer/fisson/fisson_util.h +++ b/mindspore/lite/tools/optimizer/fisson/fisson_util.h @@ -48,12 +48,12 @@ struct SplitInfo { typedef enum { CUT_N, CUT_H, CUT_W, CUT_C_IN, CUT_C_OUT, CUT_NONE } CuttingStragedy; -std::vector GetSplitPadList(const std::shared_ptr &ori_conv_prim, int64_t input_h, +std::vector GetSplitPadList(const api::SharedPtr &ori_conv_prim, int64_t input_h, int64_t input_w); bool IsConv2D(const AnfNodePtr &node); -std::shared_ptr CopyConvPrim(const std::shared_ptr &ori_attr); +api::SharedPtr CopyConvPrim(const api::SharedPtr &ori_attr); bool UpdateSplitInfo(const FuncGraphPtr &func_graph, const std::vector &conv_nodes, SplitInfo *split_info); diff --git a/mindspore/lite/tools/optimizer/fisson/iter_node_outputs.cc b/mindspore/lite/tools/optimizer/fisson/iter_node_outputs.cc index 9f46afcd77..857fb0938b 100644 --- a/mindspore/lite/tools/optimizer/fisson/iter_node_outputs.cc +++ b/mindspore/lite/tools/optimizer/fisson/iter_node_outputs.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fisson/iter_node_outputs.h" #include "tools/optimizer/parallel/spliter.h" #include "nnacl/op_base.h" diff --git a/mindspore/lite/tools/optimizer/fisson/multi_conv_split_pass.cc b/mindspore/lite/tools/optimizer/fisson/multi_conv_split_pass.cc index 3fc15d893b..46b27d964e 100644 --- a/mindspore/lite/tools/optimizer/fisson/multi_conv_split_pass.cc +++ b/mindspore/lite/tools/optimizer/fisson/multi_conv_split_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fisson/multi_conv_split_pass.h" #include #include @@ -23,6 +24,7 @@ #include "tools/optimizer/common/gllo_utils.h" #include "tools/optimizer/parallel/split_strategy.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" using mindspore::converter::FmkType; using mindspore::schema::PrimitiveType_Conv2dTransposeFusion; diff --git a/mindspore/lite/tools/optimizer/fisson/node_out_shapes.cc b/mindspore/lite/tools/optimizer/fisson/node_out_shapes.cc index ac8cb7b452..658c7b6c74 100644 --- a/mindspore/lite/tools/optimizer/fisson/node_out_shapes.cc +++ b/mindspore/lite/tools/optimizer/fisson/node_out_shapes.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fisson/node_out_shapes.h" #include #include diff --git a/mindspore/lite/tools/optimizer/format/delete_redundant_transpose.cc b/mindspore/lite/tools/optimizer/format/delete_redundant_transpose.cc index e01d5c8d9a..a80d9e2e3d 100644 --- a/mindspore/lite/tools/optimizer/format/delete_redundant_transpose.cc +++ b/mindspore/lite/tools/optimizer/format/delete_redundant_transpose.cc @@ -14,10 +14,12 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/format/delete_redundant_transpose.h" #include #include "tools/optimizer/common/format_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/format/to_format_base.cc b/mindspore/lite/tools/optimizer/format/to_format_base.cc index fd4e5b698a..f035ae2f39 100644 --- a/mindspore/lite/tools/optimizer/format/to_format_base.cc +++ b/mindspore/lite/tools/optimizer/format/to_format_base.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/format/to_format_base.h" #include #include "ops/op_utils.h" diff --git a/mindspore/lite/tools/optimizer/format/to_nchw_format.cc b/mindspore/lite/tools/optimizer/format/to_nchw_format.cc index 9d753bec7c..79e8dab5f5 100644 --- a/mindspore/lite/tools/optimizer/format/to_nchw_format.cc +++ b/mindspore/lite/tools/optimizer/format/to_nchw_format.cc @@ -14,7 +14,9 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/format/to_nchw_format.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/format/to_nhwc_format.cc b/mindspore/lite/tools/optimizer/format/to_nhwc_format.cc index af8abcbd47..79117e99df 100644 --- a/mindspore/lite/tools/optimizer/format/to_nhwc_format.cc +++ b/mindspore/lite/tools/optimizer/format/to_nhwc_format.cc @@ -14,8 +14,10 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/format/to_nhwc_format.h" #include +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/fusion/activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/activation_fusion.cc index dff5f9f51a..6ea960124a 100644 --- a/mindspore/lite/tools/optimizer/fusion/activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/activation_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/activation_fusion.h" #include #include @@ -26,9 +27,9 @@ namespace mindspore { namespace opt { STATUS DoFusion(CNodePtr cur_cnode, const CNodePtr &pre_cnode) { - auto cur_act_prim = GetValueNode>(cur_cnode->input(0)); + auto cur_act_prim = ops::GetOperator(cur_cnode->input(0)); MS_ASSERT(cur_act_prim != nullptr); - auto pre_act_prim = GetValueNode>(pre_cnode->input(0)); + auto pre_act_prim = ops::GetOperator(pre_cnode->input(0)); MS_ASSERT(pre_act_prim != nullptr); MS_CHECK_TRUE_MSG(cur_act_prim->GetAttr(ops::kActivationType) != nullptr, false, "Get activation type failed."); MS_CHECK_TRUE_MSG(pre_act_prim->GetAttr(ops::kActivationType) != nullptr, false, "Get activation type failed."); diff --git a/mindspore/lite/tools/optimizer/fusion/add_concat_activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/add_concat_activation_fusion.cc index 280e2451f2..7769636f54 100644 --- a/mindspore/lite/tools/optimizer/fusion/add_concat_activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/add_concat_activation_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/add_concat_activation_fusion.h" #include #include "ops/concat.h" @@ -21,6 +22,7 @@ #include "ops/fusion/add_fusion.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { const BaseRef AddConcatActivationFusion::DefinePattern() const { @@ -60,10 +62,10 @@ const AnfNodePtr AddConcatActivationFusion::Process(const FuncGraphPtr &func_gra return nullptr; } auto right_add_cnode = right_add_node->cast(); - auto right_add_prim = GetValueNode>(right_add_cnode->input(0)); + auto right_add_prim = ops::GetOperator(right_add_cnode->input(0)); MS_CHECK_TRUE_RET(right_add_prim != nullptr, nullptr); if (right_add_prim->GetAttr(ops::kActivationType) == nullptr) { - right_add_prim->AddAttr(ops::kActivationType, MakeValue(ActivationType::NO_ACTIVATION)); + right_add_prim->AddAttr(ops::kActivationType, api::MakeValue(ActivationType::NO_ACTIVATION)); } if (right_add_prim->get_activation_type() != ActivationType::NO_ACTIVATION) { MS_LOG(INFO) << "right add node has activation"; @@ -76,17 +78,16 @@ const AnfNodePtr AddConcatActivationFusion::Process(const FuncGraphPtr &func_gra return nullptr; } auto left_add_cnode = left_add_node->cast(); - auto left_add_prim = GetValueNode>(left_add_cnode->input(0)); + auto left_add_prim = ops::GetOperator(left_add_cnode->input(0)); MS_CHECK_TRUE_RET(left_add_prim != nullptr, nullptr); if (left_add_prim->GetAttr(ops::kActivationType) == nullptr) { - left_add_prim->AddAttr(ops::kActivationType, MakeValue(ActivationType::NO_ACTIVATION)); + left_add_prim->AddAttr(ops::kActivationType, api::MakeValue(ActivationType::NO_ACTIVATION)); } if (left_add_prim->get_activation_type() != ActivationType::NO_ACTIVATION) { MS_LOG(INFO) << "left add node has activation"; return nullptr; } - - auto act_prim = GetValueNode>(act_cnode->input(0)); + auto act_prim = ops::GetOperator(act_cnode->input(0)); MS_CHECK_TRUE_RET(act_prim != nullptr, nullptr); if (act_prim->GetAttr(ops::kActivationType) != nullptr) { right_add_prim->set_activation_type(act_prim->get_activation_type()); diff --git a/mindspore/lite/tools/optimizer/fusion/affine_activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/affine_activation_fusion.cc index d78936080d..c1e1bfa9cb 100644 --- a/mindspore/lite/tools/optimizer/fusion/affine_activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/affine_activation_fusion.cc @@ -14,12 +14,14 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/affine_activation_fusion.h" #include #include "tools/optimizer/common/gllo_utils.h" #include "ops/fusion/activation.h" #include "ops/affine.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { const BaseRef AffineActivationFusion::DefinePattern() const { @@ -51,7 +53,7 @@ const AnfNodePtr AffineActivationFusion::Process(const FuncGraphPtr &func_graph, lite::ReturnCode::GetSingleReturnCode()->UpdateReturnCode(lite::RET_NULL_PTR); return nullptr; } - auto activation_prim = GetValueNode>(activation_node->input(kAnfPrimitiveIndex)); + auto activation_prim = ops::GetOperator(activation_node->input(kAnfPrimitiveIndex)); MS_ASSERT(activation_prim != nullptr); AnfNodePtr pre_node = activation_node->input(1); if (!CheckPrimitiveType(pre_node, prim::kPrimAffine)) { @@ -68,7 +70,7 @@ const AnfNodePtr AffineActivationFusion::Process(const FuncGraphPtr &func_graph, if (IsMarkedTrainOp(affine_node)) { return nullptr; } - auto affine_prim = GetValueNode>(affine_node->input(kAnfPrimitiveIndex)); + auto affine_prim = ops::GetOperator(affine_node->input(kAnfPrimitiveIndex)); MS_ASSERT(affine_prim != nullptr); if (!activation_prim->HasAttr(ops::kActivationType)) { diff --git a/mindspore/lite/tools/optimizer/fusion/affine_fusion.cc b/mindspore/lite/tools/optimizer/fusion/affine_fusion.cc index e906541388..9b4839db88 100644 --- a/mindspore/lite/tools/optimizer/fusion/affine_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/affine_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/affine_fusion.h" #include #include @@ -23,6 +25,7 @@ #include "ops/mat_mul.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { constexpr auto kInputWithBiasNum = 4; @@ -60,8 +63,10 @@ const AnfNodePtr AffineFusion::Process(const FuncGraphPtr &func_graph, const Anf lite::ReturnCode::GetSingleReturnCode()->UpdateReturnCode(lite::RET_NULL_PTR); return nullptr; } - auto matmul_prim = GetValueNode>(matmul_node->input(kAnfPrimitiveIndex)); + auto matmul_prim = ops::GetOperator(matmul_node->input(kAnfPrimitiveIndex)); MS_CHECK_TRUE_RET(matmul_prim != nullptr, nullptr); + auto matmul_prim_c = matmul_prim->GetPrim(); + MS_CHECK_TRUE_RET(matmul_prim_c != nullptr, nullptr); // splice AnfNodePtr pre_node = matmul_node->input(1); if (!CheckPrimitiveType(pre_node, prim::kPrimSplice)) { @@ -78,7 +83,7 @@ const AnfNodePtr AffineFusion::Process(const FuncGraphPtr &func_graph, const Anf if (IsMarkedTrainOp(splice_node)) { return nullptr; } - auto splice_prim = GetValueNode>(splice_node->input(kAnfPrimitiveIndex)); + auto splice_prim = ops::GetOperator(splice_node->input(kAnfPrimitiveIndex)); MS_CHECK_TRUE_RET(splice_prim != nullptr, nullptr); /** * Affine attribute: @@ -90,6 +95,8 @@ const AnfNodePtr AffineFusion::Process(const FuncGraphPtr &func_graph, const Anf // new primitive auto affine_prim = std::make_shared(); MS_CHECK_TRUE_RET(affine_prim != nullptr, nullptr); + auto affine_prim_c = affine_prim->GetPrim(); + MS_CHECK_TRUE_RET(affine_prim_c != nullptr, nullptr); // copy splice attr to affine MS_CHECK_TRUE_RET(splice_prim->GetAttr(ops::kSpliceContext) != nullptr, nullptr); affine_prim->set_context(splice_prim->get_context()); @@ -103,7 +110,7 @@ const AnfNodePtr AffineFusion::Process(const FuncGraphPtr &func_graph, const Anf affine_prim->set_transpose_b(matmul_prim->get_transpose_b()); } // construct affine node - auto affine_value_node = NewValueNode(affine_prim); + auto affine_value_node = NewValueNode(affine_prim_c); MS_CHECK_TRUE_RET(affine_value_node != nullptr, nullptr); std::vector affine_inputs = {affine_value_node, splice_node->input(1), matmul_node->input(kInputIndexTwo)}; diff --git a/mindspore/lite/tools/optimizer/fusion/batchmatmul_fusion.cc b/mindspore/lite/tools/optimizer/fusion/batchmatmul_fusion.cc index 3c229212c8..819706ccca 100644 --- a/mindspore/lite/tools/optimizer/fusion/batchmatmul_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/batchmatmul_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/batchmatmul_fusion.h" #include #include @@ -23,6 +25,7 @@ #include "tools/optimizer/common/gllo_utils.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -100,6 +103,8 @@ std::shared_ptr BuildMatMulPrim(const CNodePtr &stack_cnode) MS_LOG(ERROR) << "new MatMul failed"; return nullptr; } + auto matmul_prim_c = matmul_cvalue->GetPrim(); + MS_CHECK_TRUE_RET(matmul_prim_c != nullptr, nullptr); std::vector jointed_quant_params; for (size_t i = 1; i < stack_cnode->inputs().size(); i++) { @@ -144,7 +149,7 @@ std::shared_ptr BuildMatMulPrim(const CNodePtr &stack_cnode) rmatmul_quant_params.emplace_back(jointed_quant_params); auto quant_params_holder = std::make_shared(rmatmul_quant_params, output_quant_params); MS_CHECK_TRUE_RET(quant_params_holder != nullptr, nullptr); - matmul_cvalue->AddAttr("quant_params", quant_params_holder); + matmul_prim_c->AddAttr("quant_params", quant_params_holder); return matmul_cvalue; } @@ -333,7 +338,7 @@ const AnfNodePtr BatchMatMulFusion::Process(const FuncGraphPtr &func_graph, cons MS_ASSERT(right_reshape_node != nullptr); auto matmul_cvalue = BuildMatMulPrim(stack_cnode); MS_CHECK_TRUE_RET(matmul_cvalue != nullptr, nullptr); - auto matmul_value_node = NewValueNode(std::shared_ptr(matmul_cvalue)); + auto matmul_value_node = NewValueNode(matmul_cvalue->GetPrim()); MS_CHECK_TRUE_RET(matmul_value_node != nullptr, nullptr); std::vector matmul_inputs = {matmul_value_node, left_matmul_input}; @@ -346,7 +351,7 @@ const AnfNodePtr BatchMatMulFusion::Process(const FuncGraphPtr &func_graph, cons MS_LOG(ERROR) << "GetRightMatmulInputParamter failed"; return node; } - auto prim_matmul = GetValueNode>(matmul_value_node); + auto prim_matmul = ops::GetOperator(matmul_value_node); MS_ASSERT(prim_matmul != nullptr); prim_matmul->set_transpose_b(true); matmul_inputs.push_back(rmatmul_paramter); @@ -382,7 +387,7 @@ const AnfNodePtr BatchMatMulFusion::Process(const FuncGraphPtr &func_graph, cons MS_CHECK_TRUE_RET(stack_cnode->abstract() != nullptr, nullptr); matmul_cnode->set_abstract(stack_cnode->abstract()->Clone()); if (right_transpose) { - auto matmul_primitive = GetValueNode>(matmul_cnode->input(0)); + auto matmul_primitive = ops::GetOperator(matmul_cnode->input(0)); matmul_primitive->set_transpose_b(true); } MS_LOG(INFO) << "stack node:" << stack_cnode->fullname_with_scope() << " batchmatmul fusion success"; diff --git a/mindspore/lite/tools/optimizer/fusion/batchnorm_to_scale_fusion.cc b/mindspore/lite/tools/optimizer/fusion/batchnorm_to_scale_fusion.cc index a55df42976..1fa8e2e9c1 100644 --- a/mindspore/lite/tools/optimizer/fusion/batchnorm_to_scale_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/batchnorm_to_scale_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/batchnorm_to_scale_fusion.h" #include #include "ops/batch_norm.h" @@ -24,6 +25,7 @@ #include "tools/common/tensor_util.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -314,7 +316,12 @@ bool BatchNormToScaleFusion::Run(const FuncGraphPtr &func_graph) { int64_t axis = input_shape_.size() == DIMENSION_4D ? -1 : 1; scale_primitive->set_axis(axis); scale_primitive->set_activation_type(ActivationType::NO_ACTIVATION); - auto scale_node = func_graph->NewCNode(scale_primitive, {cnode->input(1), new_weight_param, new_bias_param}); + auto scale_primitive_c = scale_primitive->GetPrim(); + if (scale_primitive_c == nullptr) { + MS_LOG(ERROR) << "new scale primitive_c failed"; + return false; + } + auto scale_node = func_graph->NewCNode(scale_primitive_c, {cnode->input(1), new_weight_param, new_bias_param}); scale_node->set_abstract(cnode->abstract()); (void)manager->Replace(cnode, scale_node); } diff --git a/mindspore/lite/tools/optimizer/fusion/conv_activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_activation_fusion.cc index a37ceee7e2..778eded480 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_activation_fusion.cc @@ -13,13 +13,15 @@ * See the License for the specific language governing permissions and * limitations under the License. */ - +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_activation_fusion.h" + #include + +#include "nnacl/op_base.h" #include "ops/fusion/activation.h" #include "ops/op_utils.h" #include "tools/optimizer/common/gllo_utils.h" -#include "nnacl/op_base.h" namespace mindspore::opt { const BaseRef ConvActivationFusion::DefinePattern() const { @@ -37,16 +39,15 @@ const AnfNodePtr ConvActivationFusion::Process(const FuncGraphPtr &func_graph, c return nullptr; } auto act_node = node->cast(); - if (IsMarkedTrainOp(act_node)) { - return nullptr; - } + MS_CHECK_TRUE_RET(IsMarkedTrainOp(act_node) != true, nullptr); if (act_node == nullptr || act_node->size() != kInputSizeTwo || !CheckPrimitiveType(act_node, prim::kPrimActivation)) { return nullptr; } - auto act_prim = GetValueNode>(act_node->input(0)); + auto act_prim = ops::GetOperator(act_node->input(0)); MS_CHECK_TRUE_MSG(act_prim != nullptr, nullptr, "activation prim is nullptr."); - if (act_prim->GetAttr(ops::kActivationType) != nullptr && act_prim->get_activation_type() != mindspore::RELU && + auto act_prim_c = act_prim->GetPrim(); + if (act_prim_c->GetAttr(ops::kActivationType) != nullptr && act_prim->get_activation_type() != mindspore::RELU && act_prim->get_activation_type() != mindspore::RELU6) { return nullptr; } diff --git a/mindspore/lite/tools/optimizer/fusion/conv_biasadd_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_biasadd_fusion.cc index 5394639bb5..63618e14dc 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_biasadd_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_biasadd_fusion.cc @@ -13,6 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_biasadd_fusion.h" #include #include @@ -21,6 +22,8 @@ #include "tools/optimizer/common/gllo_utils.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" +#include "mindapi/base/types.h" namespace mindspore::opt { namespace { diff --git a/mindspore/lite/tools/optimizer/fusion/conv_bn_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_bn_fusion.cc index e09539a914..9c0c91b454 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_bn_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_bn_fusion.cc @@ -13,7 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ - +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_bn_fusion.h" #include #include "include/common/utils/utils.h" diff --git a/mindspore/lite/tools/optimizer/fusion/conv_conv_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_conv_fusion.cc index 7a6e04232e..2b0a8ad023 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_conv_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_conv_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_conv_fusion.h" #include #include @@ -21,6 +22,7 @@ #include "ops/fusion/conv2d_fusion.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -39,12 +41,12 @@ bool IsCommonConvNode(const BaseRef &n) { if (!CheckPrimitiveType(anf_node, prim::kPrimConv2DFusion)) { return false; } - std::shared_ptr conv = nullptr; + api::SharedPtr conv = nullptr; if (utils::isa(anf_node)) { auto c_node = anf_node->cast(); - conv = GetValueNode>(c_node->input(0)); + conv = ops::GetOperator(c_node->input(0)); } else if (utils::isa(anf_node)) { - conv = GetValueNode>(anf_node); + conv = ops::GetOperator(anf_node); } if (conv == nullptr) { return false; @@ -215,9 +217,9 @@ STATUS ReplaceParametersAndNodes(const FuncGraphPtr &func_graph, const CNodePtr bool IsPrimitiveProper(const CNodePtr &up_conv_cnode, const CNodePtr &down_conv_cnode) { MS_ASSERT(up_conv_cnode != nullptr && down_conv_cnode != nullptr); - auto down_conv_primitive = GetValueNode>(down_conv_cnode->input(0)); + auto down_conv_primitive = ops::GetOperator(down_conv_cnode->input(0)); MS_ASSERT(down_conv_primitive != nullptr); - auto up_conv_primitive = GetValueNode>(up_conv_cnode->input(0)); + auto up_conv_primitive = ops::GetOperator(up_conv_cnode->input(0)); MS_ASSERT(up_conv_primitive != nullptr); int64_t up_conv_group = up_conv_primitive->GetAttr(ops::kGroup) == nullptr ? 1 : up_conv_primitive->get_group(); int64_t down_conv_group = down_conv_primitive->GetAttr(ops::kGroup) == nullptr ? 1 : down_conv_primitive->get_group(); diff --git a/mindspore/lite/tools/optimizer/fusion/conv_pad_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_pad_fusion.cc index ff7cc5686f..67e2d3d972 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_pad_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_pad_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_pad_fusion.h" #include #include @@ -22,6 +23,8 @@ #include "ops/fusion/conv2d_fusion.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/primitive_c.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -64,7 +67,7 @@ void ReplaceParamsAndNodes(const FuncGraphPtr &func_graph, const CNodePtr &conv_ pad_list_data.push_back(pad_data[kRight + NCHWTopPadPos]); } - auto conv_primitive = GetValueNode>(conv_cnode->input(0)); + auto conv_primitive = ops::GetOperator(conv_cnode->input(0)); MS_ASSERT(conv_primitive != nullptr); int64_t conv_pad_mode = conv_primitive->GetAttr(ops::kPadMode) == nullptr ? 0 : conv_primitive->get_pad_mode(); if (conv_pad_mode == PadMode::PAD) { @@ -79,7 +82,7 @@ void ReplaceParamsAndNodes(const FuncGraphPtr &func_graph, const CNodePtr &conv_ } } } else if (conv_pad_mode == PadMode::SAME) { - ValuePtr kernel_node = conv_primitive->GetAttr(ops::kKernelSize); + auto kernel_node = conv_primitive->GetAttr(ops::kKernelSize); MS_ASSERT(kernel_node != nullptr); std::vector kernel_list = GetValue>(kernel_node); if (kernel_list.size() != kFilterDimsSize) { @@ -130,16 +133,18 @@ bool IsPrimitiveProper(const CNodePtr &pad_cnode) { if (tensor->data_c() == nullptr || tensor->ElementsNum() != kPadElementNum) { return false; } - auto pad_primitive = GetValueNode>(pad_cnode->input(0)); - MS_ASSERT(pad_primitive != nullptr); - if (!pad_primitive->HasAttr(ops::kPaddingMode)) { + auto prim = GetValueNode(pad_cnode->input(0)); + MS_ASSERT(prim != nullptr); + auto pad_primitive = api::MakeShared(prim); + if (!prim->HasAttr(ops::kPaddingMode)) { return false; } + MS_ASSERT(pad_primitive != nullptr); int64_t pad_mode = pad_primitive->get_padding_mode(); if (pad_mode != PaddingMode::CONSTANT) { return false; } - ValuePtr pad_constant_node = pad_primitive->GetAttr(ops::kConstantValue); + auto pad_constant_node = pad_primitive->GetAttr(ops::kConstantValue); if (pad_constant_node == nullptr) { return false; } diff --git a/mindspore/lite/tools/optimizer/fusion/conv_scale_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_scale_fusion.cc index 371c3f97bd..399f39eb6b 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_scale_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_scale_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_scale_fusion.h" #include #include "tools/optimizer/common/gllo_utils.h" diff --git a/mindspore/lite/tools/optimizer/fusion/conv_transform_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_transform_fusion.cc index bc3a5460e2..842f78c88c 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_transform_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_transform_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_transform_fusion.h" #include #include @@ -25,6 +26,7 @@ #include "tools/converter/quant_param_holder.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -37,16 +39,20 @@ int64_t GetOutChannels(const CNodePtr &conv_node) { auto value_node = conv_node->input(0); MS_ASSERT(value_node != nullptr); if (CheckPrimitiveType(conv_node, prim::kPrimConv2DFusion)) { - auto conv_prim = GetValueNode>(value_node); + auto conv_prim = ops::GetOperator(value_node); MS_ASSERT(conv_prim != nullptr); - if (conv_prim->GetAttr(ops::kOutChannel) == nullptr) { + auto conv_prim_c = conv_prim->GetPrim(); + MS_ASSERT(conv_prim_c != nullptr); + if (conv_prim_c->GetAttr(ops::kOutChannel) == nullptr) { return 0; } return conv_prim->get_out_channel(); } else if (CheckPrimitiveType(conv_node, prim::kPrimConv2dTransposeFusion)) { - auto conv_prim = GetValueNode>(value_node); + auto conv_prim = ops::GetOperator(value_node); MS_ASSERT(conv_prim != nullptr); - if (conv_prim->GetAttr(ops::kOutChannel) == nullptr) { + auto conv_prim_c = conv_prim->GetPrim(); + MS_ASSERT(conv_prim_c != nullptr); + if (conv_prim_c->GetAttr(ops::kOutChannel) == nullptr) { return 0; } return conv_prim->get_out_channel(); @@ -322,9 +328,11 @@ int ConvTransformFusion::CalNewWeightTensor(const CNodePtr &conv_node, const ten if (CheckPrimitiveType(conv_node, prim::kPrimConv2DFusion)) { GenerateNewWeightConv2D(tmp_weight_data, weight_data, trans_scale, weight_shape_size, kernel_num); } else if (CheckPrimitiveType(conv_node, prim::kPrimConv2dTransposeFusion) && !is_depth_wise) { - auto conv_primc = conv_prim->cast>(); - MS_ASSERT(conv_primc != nullptr); - auto group = conv_primc->GetAttr(ops::kGroup) == nullptr ? 1 : conv_primc->get_group(); + auto conv2d_prim = api::MakeShared(conv_prim); + MS_ASSERT(conv2d_prim != nullptr); + auto conv2d_prim_c = conv2d_prim->GetPrim(); + MS_ASSERT(conv2d_prim_c != nullptr); + auto group = conv2d_prim_c->GetAttr(ops::kGroup) == nullptr ? 1 : conv2d_prim->get_group(); GenerateNewWeightConv2DTranspose(tmp_weight_data, trans_scale, weight_tensor, group, kernel_num); } auto ret = memcpy_s(weight_data, weight_tensor->Size(), tmp_weight_data, data_size); diff --git a/mindspore/lite/tools/optimizer/fusion/conv_tuple_activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_tuple_activation_fusion.cc index 25c8515229..f9b8e0b280 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_tuple_activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_tuple_activation_fusion.cc @@ -14,12 +14,14 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_tuple_activation_fusion.h" #include #include "ops/fusion/activation.h" #include "ops/fusion/conv2d_fusion.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { const BaseRef ConvTupleActivationFusion::DefinePattern() const { @@ -50,9 +52,11 @@ const AnfNodePtr ConvTupleActivationFusion::Process(const FuncGraphPtr &func_gra if (!CheckPrimitiveType(act_node, prim::kPrimActivation)) { return nullptr; } - auto act_prim = GetValueNode>(act_node->input(0)); + auto act_prim = ops::GetOperator(act_node->input(0)); MS_ASSERT(act_prim != nullptr); - if (act_prim->GetAttr(ops::kActivationType) == nullptr || + auto act_prim_c = act_prim->GetPrim(); + MS_ASSERT(act_prim_c != nullptr); + if (act_prim_c->GetAttr(ops::kActivationType) == nullptr || (act_prim->get_activation_type() != mindspore::RELU && act_prim->get_activation_type() != mindspore::RELU6)) { return nullptr; } @@ -73,10 +77,13 @@ const AnfNodePtr ConvTupleActivationFusion::Process(const FuncGraphPtr &func_gra return nullptr; } if (CheckPrimitiveType(conv_node, prim::kPrimConv2DFusion)) { - auto primc = GetValueNode>(conv_cnode->input(0)); - MS_ASSERT(primc != nullptr); - if (primc->GetAttr(ops::kActivationType) == nullptr || primc->get_activation_type() == mindspore::NO_ACTIVATION) { - primc->set_activation_type(act_prim->get_activation_type()); + auto conv_prim = ops::GetOperator(conv_cnode->input(0)); + MS_ASSERT(conv_prim != nullptr); + auto conv_prim_c = conv_prim->GetPrim(); + MS_ASSERT(conv_prim_c != nullptr); + if (conv_prim_c->GetAttr(ops::kActivationType) == nullptr || + conv_prim->get_activation_type() == mindspore::NO_ACTIVATION) { + conv_prim->set_activation_type(act_prim->get_activation_type()); conv_node->set_abstract(act_node->abstract()); return conv_node; } diff --git a/mindspore/lite/tools/optimizer/fusion/conv_tuplegetitem_fusion.cc b/mindspore/lite/tools/optimizer/fusion/conv_tuplegetitem_fusion.cc index 7b22608dfe..44baf41cc0 100644 --- a/mindspore/lite/tools/optimizer/fusion/conv_tuplegetitem_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/conv_tuplegetitem_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/conv_tuplegetitem_fusion.h" #include #include "tools/optimizer/common/gllo_utils.h" diff --git a/mindspore/lite/tools/optimizer/fusion/fullconnected_add_fusion.cc b/mindspore/lite/tools/optimizer/fusion/fullconnected_add_fusion.cc index 646f0d7dee..ce5be51a78 100644 --- a/mindspore/lite/tools/optimizer/fusion/fullconnected_add_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/fullconnected_add_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/fullconnected_add_fusion.h" #include #include @@ -21,6 +22,7 @@ #include "ops/fusion/full_connection.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -58,14 +60,16 @@ bool IsPrimitiveProper(const CNodePtr &add_cnode, const CNodePtr &fc_cnode, int return false; } } - auto fc_primc = GetValueNode>(fc_cnode->input(0)); + auto fc_primc = ops::GetOperator(fc_cnode->input(0)); MS_CHECK_TRUE_RET(fc_primc != nullptr, false); if (fc_primc->GetAttr(ops::kActivationType) != nullptr && fc_primc->get_activation_type() != ActivationType::NO_ACTIVATION) { MS_LOG(INFO) << fc_cnode->fullname_with_scope() << " has activation attr"; return false; } - if (IsQuantParameterNode(fc_primc)) { + auto prim_c = fc_primc->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, false); + if (IsQuantParameterNode(prim_c)) { MS_LOG(INFO) << fc_cnode->fullname_with_scope() << "is quant node"; return false; } @@ -180,11 +184,11 @@ AnfNodePtr FullconnectedAddFusion::Process(const std::string &pattern_name, cons } if (CheckPrimitiveType(node, prim::kPrimAddFusion)) { - auto add_primc = GetValueNode>(add_cnode->input(0)); + auto add_primc = ops::GetOperator(add_cnode->input(0)); MS_CHECK_TRUE_RET(add_primc != nullptr, nullptr); if (add_primc->GetAttr(ops::kActivationType) != nullptr && add_primc->get_activation_type() != ActivationType::NO_ACTIVATION) { - auto fc_primc = GetValueNode>(fc_cnode->input(0)); + auto fc_primc = ops::GetOperator(fc_cnode->input(0)); MS_CHECK_TRUE_RET(fc_primc != nullptr, nullptr); fc_primc->set_activation_type(add_primc->get_activation_type()); } diff --git a/mindspore/lite/tools/optimizer/fusion/fullconnected_fusion.cc b/mindspore/lite/tools/optimizer/fusion/fullconnected_fusion.cc index 038421d669..20d5078aa4 100644 --- a/mindspore/lite/tools/optimizer/fusion/fullconnected_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/fullconnected_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/fullconnected_fusion.h" #include #include @@ -22,8 +23,10 @@ #include "tools/optimizer/common/gllo_utils.h" #include "tools/converter/quant_param_holder.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { + namespace { constexpr size_t kFcWeightIndex = 2; constexpr size_t kFcParameterDims = 2; @@ -190,10 +193,12 @@ bool IsPrimitiveProper(const CNodePtr &curr_fc_cnode, const CNodePtr &prev_fc_cn MS_LOG(INFO) << pre_fc_weight_node->fullname_with_scope() << "'s weight is not parameter"; return false; } - auto primc = utils::cast>(prev_primc); - MS_ASSERT(primc != nullptr); - if (primc->GetAttr(ops::kActivationType) != nullptr) { - auto activate_type = primc->get_activation_type(); + auto full_prim = api::MakeShared(prev_primc); + MS_ASSERT(full_prim != nullptr); + auto full_prim_c = full_prim->GetPrim(); + MS_ASSERT(full_prim_c != nullptr); + if (full_prim_c->GetAttr(ops::kActivationType) != nullptr) { + auto activate_type = full_prim->get_activation_type(); if (activate_type != NO_ACTIVATION) { MS_LOG(INFO) << pre_fc_weight_node->fullname_with_scope() << " has activation operator"; return false; diff --git a/mindspore/lite/tools/optimizer/fusion/gelu_fusion.cc b/mindspore/lite/tools/optimizer/fusion/gelu_fusion.cc index e3598fb79d..83210dc0f8 100644 --- a/mindspore/lite/tools/optimizer/fusion/gelu_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/gelu_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/gelu_fusion.h" #include #include @@ -34,11 +35,13 @@ CNodePtr GeLUFusion::CreateGeLUNode(const FuncGraphPtr &func_graph, const AnfNod MS_ASSERT(func_graph != nullptr && node != nullptr && equiv != nullptr); auto gelu_prim = std::make_shared(); MS_CHECK_TRUE_RET(gelu_prim != nullptr, nullptr); + auto gelu_prim_c = gelu_prim->GetPrim(); + MS_CHECK_TRUE_RET(gelu_prim_c != nullptr, nullptr); gelu_prim->set_activation_type(mindspore::GELU); gelu_prim->set_approximate(approximate_); auto input_node = utils::cast((*equiv)[input_]); MS_ASSERT(input_node != nullptr); - auto gelu_cnode = func_graph->NewCNode(gelu_prim, {input_node}); + auto gelu_cnode = func_graph->NewCNode(gelu_prim_c, {input_node}); MS_CHECK_TRUE_RET(gelu_cnode != nullptr, nullptr); gelu_cnode->set_fullname_with_scope(node->fullname_with_scope() + "_gelu"); if (node->abstract() != nullptr) { diff --git a/mindspore/lite/tools/optimizer/fusion/glu_fusion.cc b/mindspore/lite/tools/optimizer/fusion/glu_fusion.cc index 47616a2627..d229c35367 100644 --- a/mindspore/lite/tools/optimizer/fusion/glu_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/glu_fusion.cc @@ -13,6 +13,7 @@ * See the License for the specific language governing permissions and * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/glu_fusion.h" #include #include @@ -27,6 +28,8 @@ CNodePtr GLUFusion::CreateGLUNode(const FuncGraphPtr &func_graph, const AnfNodeP MS_ASSERT(func_graph != nullptr && node != nullptr && equiv != nullptr); auto glu_prim = std::make_shared(); MS_CHECK_TRUE_RET(glu_prim != nullptr, nullptr); + auto glu_prim_c = glu_prim->GetPrim(); + MS_CHECK_TRUE_RET(glu_prim_c != nullptr, nullptr); auto split_prim = GetValueNode(utils::cast((*equiv)[split_prim_])); if (split_prim != nullptr && split_prim->GetAttr(ops::kAxis) != nullptr) { auto axis = GetValue(split_prim->GetAttr(ops::kAxis)); @@ -34,7 +37,7 @@ CNodePtr GLUFusion::CreateGLUNode(const FuncGraphPtr &func_graph, const AnfNodeP } auto input_node = utils::cast((*equiv)[input_]); MS_ASSERT(input_node != nullptr); - auto glu_cnode = func_graph->NewCNode(glu_prim, {input_node}); + auto glu_cnode = func_graph->NewCNode(glu_prim_c, {input_node}); MS_CHECK_TRUE_RET(glu_cnode != nullptr, nullptr); glu_cnode->set_fullname_with_scope(node->fullname_with_scope() + "_glu"); if (node->abstract() != nullptr) { diff --git a/mindspore/lite/tools/optimizer/fusion/matmul_activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/matmul_activation_fusion.cc index 1e3dde0a78..75a3126233 100644 --- a/mindspore/lite/tools/optimizer/fusion/matmul_activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/matmul_activation_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/matmul_activation_fusion.h" #include #include "ops/fusion/activation.h" @@ -21,6 +22,7 @@ #include "include/common/utils/utils.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { const BaseRef MatMulActivationFusion::DefinePattern() const { @@ -47,7 +49,7 @@ const AnfNodePtr MatMulActivationFusion::Process(const FuncGraphPtr &func_graph, } MS_CHECK_TRUE_RET(act_cnode->input(1) != nullptr, nullptr); auto matmul_cnode = act_cnode->input(1)->cast(); - auto matmul_prim = GetValueNode>(matmul_cnode->input(0)); + auto matmul_prim = ops::GetOperator(matmul_cnode->input(0)); if (matmul_prim == nullptr) { MS_LOG(ERROR) << "matmul prim is nullptr."; return nullptr; @@ -57,7 +59,7 @@ const AnfNodePtr MatMulActivationFusion::Process(const FuncGraphPtr &func_graph, MS_LOG(ERROR) << "matmul has activation."; return nullptr; } - auto act_prim = GetValueNode>(act_cnode->input(0)); + auto act_prim = ops::GetOperator(act_cnode->input(0)); if (act_prim == nullptr) { MS_LOG(ERROR) << "activation prim is nullptr."; return nullptr; @@ -71,7 +73,7 @@ const AnfNodePtr MatMulActivationFusion::Process(const FuncGraphPtr &func_graph, if (type != mindspore::RELU && type != RELU6) { return nullptr; } - matmul_prim->AddAttr(ops::kActivationType, MakeValue(type)); + matmul_prim->AddAttr(ops::kActivationType, api::MakeValue(type)); auto manage = Manage(func_graph); if (manage == nullptr) { MS_LOG(ERROR) << "manage is nullptr."; diff --git a/mindspore/lite/tools/optimizer/fusion/matmul_add_fusion.cc b/mindspore/lite/tools/optimizer/fusion/matmul_add_fusion.cc index 9b258ba464..b67bf678a3 100644 --- a/mindspore/lite/tools/optimizer/fusion/matmul_add_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/matmul_add_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/matmul_add_fusion.h" #include #include @@ -21,6 +22,7 @@ #include "ops/fusion/mat_mul_fusion.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -58,14 +60,16 @@ bool IsPrimitiveProper(const CNodePtr &add_cnode, const CNodePtr &matmul_cnode, return false; } } - auto matmul_primc = GetValueNode>(matmul_cnode->input(0)); + auto matmul_primc = ops::GetOperator(matmul_cnode->input(0)); MS_CHECK_TRUE_RET(matmul_primc != nullptr, false); if (matmul_primc->GetAttr(ops::kActivationType) != nullptr && matmul_primc->get_activation_type() != ActivationType::NO_ACTIVATION) { MS_LOG(INFO) << matmul_cnode->fullname_with_scope() << " has activation attr"; return false; } - if (IsQuantParameterNode(matmul_primc)) { + auto matmul_prim_c = matmul_primc->GetPrim(); + MS_CHECK_TRUE_RET(matmul_prim_c != nullptr, false); + if (IsQuantParameterNode(matmul_prim_c)) { MS_LOG(INFO) << matmul_cnode->fullname_with_scope() << "is quant node"; return false; } @@ -154,11 +158,11 @@ bool MatMulAddFusion::Run(const FuncGraphPtr &func_graph) { } if (CheckPrimitiveType(node, prim::kPrimAddFusion)) { - auto add_primc = GetValueNode>(add_cnode->input(0)); + auto add_primc = ops::GetOperator(add_cnode->input(0)); MS_CHECK_TRUE_RET(add_primc != nullptr, false); if (add_primc->GetAttr(ops::kActivationType) != nullptr && add_primc->get_activation_type() != ActivationType::NO_ACTIVATION) { - auto matmul_primc = GetValueNode>(matmul_cnode->input(0)); + auto matmul_primc = ops::GetOperator(matmul_cnode->input(0)); MS_CHECK_TRUE_RET(matmul_primc != nullptr, false); matmul_primc->set_activation_type(add_primc->get_activation_type()); } diff --git a/mindspore/lite/tools/optimizer/fusion/matmul_mul_fusion.cc b/mindspore/lite/tools/optimizer/fusion/matmul_mul_fusion.cc index 91e66b0a19..3664606065 100644 --- a/mindspore/lite/tools/optimizer/fusion/matmul_mul_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/matmul_mul_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/matmul_mul_fusion.h" #include #include @@ -21,6 +22,7 @@ #include "ops/fusion/mul_fusion.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -50,7 +52,7 @@ int CalNewCnodeScale(const CNodePtr &mul_cnode, const CNodePtr &matmul_cnode) { auto matmul_weight_data = reinterpret_cast(matmul_weight_tensor->data_c()); MS_CHECK_TRUE_RET(matmul_weight_data != nullptr, RET_ERROR); - auto matmul_prim = GetValueNode>(matmul_cnode->input(0)); + auto matmul_prim = ops::GetOperator(matmul_cnode->input(0)); MS_CHECK_TRUE_RET(matmul_prim->GetAttr(ops::kTransposeB) != nullptr, RET_ERROR); bool transpose_b = matmul_prim->get_transpose_b(); @@ -141,15 +143,16 @@ bool IsPrimitiveProper(const CNodePtr &mul_cnode, const CNodePtr &matmul_cnode) } } - auto matmul_primc = GetValueNode>(matmul_cnode->input(0)); + auto matmul_primc = ops::GetOperator(matmul_cnode->input(0)); MS_CHECK_TRUE_RET(matmul_primc != nullptr, false); if (matmul_primc->GetAttr(ops::kActivationType) != nullptr && matmul_primc->get_activation_type() != ActivationType::NO_ACTIVATION) { MS_LOG(INFO) << matmul_cnode->fullname_with_scope() << " has activation attr"; return false; } - MS_CHECK_TRUE_RET(matmul_primc != nullptr, false); - if (IsQuantParameterNode(matmul_primc)) { + auto matmul_prim_c = matmul_primc->GetPrim(); + MS_CHECK_TRUE_RET(matmul_prim_c != nullptr, false); + if (IsQuantParameterNode(matmul_prim_c)) { MS_LOG(INFO) << matmul_cnode->fullname_with_scope() << "is quant node"; return false; } @@ -165,7 +168,7 @@ bool IsPrimitiveProper(const CNodePtr &mul_cnode, const CNodePtr &matmul_cnode) MS_CHECK_TRUE_RET(matmul_weight_tensor != nullptr, RET_ERROR); std::vector matmul_weight_shape = matmul_weight_tensor->shape(); MS_CHECK_TRUE_RET(matmul_weight_shape.size() >= KMatmulWeightDims, false); - auto matmul_prim = GetValueNode>(matmul_cnode->input(0)); + auto matmul_prim = ops::GetOperator(matmul_cnode->input(0)); MS_CHECK_TRUE_RET(matmul_prim->GetAttr(ops::kTransposeB) != nullptr, false); int64_t last_dim_size = matmul_prim->get_transpose_b() ? matmul_weight_shape[matmul_weight_shape.size() - kSecondToLastDim] @@ -226,9 +229,9 @@ const AnfNodePtr MatMulMulFusion::Process(const FuncGraphPtr &func_graph, const return nullptr; } - auto mul_primc = GetValueNode>(mul_cnode->input(0)); + auto mul_primc = ops::GetOperator(mul_cnode->input(0)); MS_CHECK_TRUE_RET(mul_primc != nullptr, nullptr); - auto matmul_primc = GetValueNode>(matmul_cnode->input(0)); + auto matmul_primc = ops::GetOperator(matmul_cnode->input(0)); MS_CHECK_TRUE_RET(matmul_primc != nullptr, nullptr); if (mul_primc->GetAttr(ops::kActivationType) != nullptr && mul_primc->get_activation_type() != ActivationType::NO_ACTIVATION) { diff --git a/mindspore/lite/tools/optimizer/fusion/matmul_scale_fusion.cc b/mindspore/lite/tools/optimizer/fusion/matmul_scale_fusion.cc index 933da85fb2..422b54e351 100644 --- a/mindspore/lite/tools/optimizer/fusion/matmul_scale_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/matmul_scale_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/matmul_scale_fusion.h" #include #include @@ -21,6 +22,7 @@ #include "tools/optimizer/common/gllo_utils.h" #include "tools/converter/quant_param_holder.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -50,7 +52,7 @@ int MatMulScaleFusion::CalNewBiasImpl(float *curr_weight_data, float *curr_bias_ int MatMulScaleFusion::CalNewScaleImpl(float *curr_weight_data, std::vector prev_weight_shape, float *prev_weight_data, const AnfNodePtr &prim) const { - auto matmul_prim = GetValueNode>(prim); + auto matmul_prim = ops::GetOperator(prim); auto trans_attr = matmul_prim->GetAttr(ops::kTransposeB); MS_CHECK_TRUE_RET(trans_attr != nullptr, RET_ERROR); bool transpose_b = matmul_prim->get_transpose_b(); diff --git a/mindspore/lite/tools/optimizer/fusion/mul_add_fusion.cc b/mindspore/lite/tools/optimizer/fusion/mul_add_fusion.cc index d9da5a97a6..e3f1e60ddf 100644 --- a/mindspore/lite/tools/optimizer/fusion/mul_add_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/mul_add_fusion.cc @@ -14,17 +14,18 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/mul_add_fusion.h" #include #include #include -#include "ops/fusion/mul_fusion.h" +#include "nnacl/op_base.h" #include "ops/fusion/add_fusion.h" +#include "ops/fusion/mul_fusion.h" #include "ops/fusion/scale_fusion.h" #include "ops/op_utils.h" -#include "tools/optimizer/common/gllo_utils.h" #include "tools/anf_exporter/fetch_content.h" -#include "nnacl/op_base.h" +#include "tools/optimizer/common/gllo_utils.h" namespace mindspore::opt { VectorRef MulAddFusion::DefineMulFirstPattern() const { @@ -63,9 +64,11 @@ bool MulAddFusion::CheckAddNode(const mindspore::CNodePtr &cnode) const { if (IsMarkedTrainOp(cnode)) { return false; } - auto add_primitive = GetValueNode>(cnode->input(0)); + auto add_primitive = ops::GetOperator(cnode->input(0)); MS_CHECK_TRUE_RET(add_primitive != nullptr, false); - auto quant_attr = add_primitive->GetAttr("quant_params"); + auto add_primitive_c = add_primitive->GetPrim(); + MS_CHECK_TRUE_RET(add_primitive_c != nullptr, false); + auto quant_attr = add_primitive_c->GetAttr("quant_params"); if (quant_attr != nullptr) { auto quant_param_holder = quant_attr->cast(); MS_CHECK_TRUE_RET(quant_param_holder != nullptr, false); @@ -79,7 +82,7 @@ bool MulAddFusion::CheckAddNode(const mindspore::CNodePtr &cnode) const { } ActivationType add_act_type = ActivationType::NO_ACTIVATION; - if (add_primitive->GetAttr(ops::kActivationType) != nullptr) { + if (add_primitive_c->GetAttr(ops::kActivationType) != nullptr) { add_act_type = add_primitive->get_activation_type(); if (add_act_type != ActivationType::RELU && add_act_type != ActivationType::RELU6 && add_act_type != ActivationType::NO_ACTIVATION) { @@ -100,9 +103,11 @@ bool MulAddFusion::CheckMulNode(const mindspore::FuncGraphPtr &func_graph, const if (IsMarkedTrainOp(cnode)) { return false; } - auto mul_primitive = GetValueNode>(cnode->input(0)); + auto mul_primitive = ops::GetOperator(cnode->input(0)); MS_CHECK_TRUE_RET(mul_primitive != nullptr, false); - auto quant_attr = mul_primitive->GetAttr("quant_params"); + auto mul_primitive_c = mul_primitive->GetPrim(); + MS_CHECK_TRUE_RET(mul_primitive_c != nullptr, false); + auto quant_attr = mul_primitive_c->GetAttr("quant_params"); if (quant_attr != nullptr) { auto quant_param_holder = quant_attr->cast(); MS_CHECK_TRUE_RET(quant_param_holder != nullptr, false); @@ -115,7 +120,7 @@ bool MulAddFusion::CheckMulNode(const mindspore::FuncGraphPtr &func_graph, const } } - if (mul_primitive->GetAttr(ops::kActivationType) != nullptr && + if (mul_primitive_c->GetAttr(ops::kActivationType) != nullptr && mul_primitive->get_activation_type() != ActivationType::NO_ACTIVATION) { MS_LOG(DEBUG) << "Only support mul node with no activation"; return false; @@ -243,10 +248,12 @@ AnfNodePtr MulAddFusion::Process(const std::string &pattern_name, const mindspor return nullptr; } scale_primitive->set_activation_type(scale_act_type_); + auto scale_primitive_c = scale_primitive->GetPrim(); + MS_CHECK_TRUE_RET(scale_primitive_c != nullptr, nullptr); scale_primitive->set_axis(-(static_cast(bias_tensor_->shape_c().size() + axis_offset))); // create scale op - auto scale_node = func_graph->NewCNode(scale_primitive, {mul_input_anode, mul_const_anode_, add_const_anode_}); + auto scale_node = func_graph->NewCNode(scale_primitive_c, {mul_input_anode, mul_const_anode_, add_const_anode_}); scale_node->set_abstract(add_cnode->abstract()); return scale_node; } diff --git a/mindspore/lite/tools/optimizer/fusion/multi_head_attention_fusion.cc b/mindspore/lite/tools/optimizer/fusion/multi_head_attention_fusion.cc index c68ea8ee90..f1623afac5 100644 --- a/mindspore/lite/tools/optimizer/fusion/multi_head_attention_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/multi_head_attention_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/multi_head_attention_fusion.h" #include #include @@ -262,7 +264,9 @@ CNodePtr MultiHeadAttentionFusion::CreateMultiHeadAttentionNode(const FuncGraphP MS_LOG(ERROR) << "Build attention primitive failed."; return nullptr; } - auto value_node = NewValueNode(attention_prim); + auto attention_prim_c = attention_prim->GetPrim(); + MS_CHECK_TRUE_RET(attention_prim_c != nullptr, nullptr); + auto value_node = NewValueNode(attention_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto input_q = utils::cast((*equiv)[input_q_]); auto input_k = utils::cast((*equiv)[input_k_]); @@ -358,7 +362,9 @@ CNodePtr MultiHeadAttentionFusion::CreateMaskedMultiHeadAttentionNode(const Func MS_LOG(ERROR) << "Build attention primitive failed."; return nullptr; } - auto value_node = NewValueNode(attention_prim); + auto attention_prim_c = attention_prim->GetPrim(); + MS_CHECK_TRUE_RET(attention_prim_c != nullptr, nullptr); + auto value_node = NewValueNode(attention_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto input_q = utils::cast((*equiv)[input_q_]); auto input_k = utils::cast((*equiv)[input_k_]); diff --git a/mindspore/lite/tools/optimizer/fusion/norm_fusion.cc b/mindspore/lite/tools/optimizer/fusion/norm_fusion.cc index c63acff4ba..06362b9b42 100644 --- a/mindspore/lite/tools/optimizer/fusion/norm_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/norm_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/norm_fusion.h" #include #include @@ -24,6 +26,7 @@ #include "securec/include/securec.h" #include "nnacl/op_base.h" #include "src/ops/ops_utils.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -70,9 +73,10 @@ bool IsReduceNode(const EquivPtr &equiv, const VarPtr &input_prim, const VarPtr MS_ASSERT(equiv != nullptr && input_prim != nullptr && input_axes != nullptr && axes != nullptr); auto reduce_value = utils::cast((*equiv)[input_prim]); MS_ASSERT(reduce_value != nullptr); - auto mean2_primitive = GetValueNode>(reduce_value); - if (mean2_primitive == nullptr || mean2_primitive->GetAttr(ops::kMode) == nullptr || - mean2_primitive->get_mode() != mindspore::Reduce_Mean) { + auto mean2_primitive = ops::GetOperator(reduce_value); + MS_CHECK_TRUE_RET(mean2_primitive != nullptr, false); + auto mean2_primitive_c = mean2_primitive->GetPrim(); + if (mean2_primitive_c->GetAttr(ops::kMode) == nullptr || mean2_primitive->get_mode() != mindspore::Reduce_Mean) { return false; } if (GetReduceAxes((*equiv)[input_axes], axes) != lite::RET_OK) { @@ -107,21 +111,25 @@ CNodePtr NormFusion::CreateNormNode(const FuncGraphPtr &func_graph, const EquivP int begin_params_axis) const { MS_ASSERT(func_graph != nullptr); MS_ASSERT(equiv != nullptr); - PrimitiveCPtr primitive = nullptr; + PrimitiveCPtr primitive_c = nullptr; if (type == schema::PrimitiveType_LayerNormFusion) { auto layer_norm_primitive = std::make_shared(); MS_CHECK_TRUE_RET(layer_norm_primitive != nullptr, nullptr); layer_norm_primitive->Init(begin_norm_axis, begin_params_axis, epsilon, true); - primitive = layer_norm_primitive; + auto layer_norm_primitive_c = layer_norm_primitive->GetPrim(); + MS_CHECK_TRUE_RET(layer_norm_primitive_c != nullptr, nullptr); + primitive_c = layer_norm_primitive_c; } else if (type == schema::PrimitiveType_InstanceNorm) { auto instance_norm_primitive = std::make_shared(); MS_CHECK_TRUE_RET(instance_norm_primitive != nullptr, nullptr); + auto instance_norm_primitive_c = instance_norm_primitive->GetPrim(); + MS_CHECK_TRUE_RET(instance_norm_primitive_c != nullptr, nullptr); instance_norm_primitive->Init(epsilon); - primitive = instance_norm_primitive; + primitive_c = instance_norm_primitive_c; } else { return nullptr; } - auto value_node = NewValueNode(primitive); + auto value_node = NewValueNode(primitive_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); std::vector new_node_inputs = {value_node}; auto input_node = utils::cast((*equiv)[input_]); diff --git a/mindspore/lite/tools/optimizer/fusion/onnx_gelu_fusion.cc b/mindspore/lite/tools/optimizer/fusion/onnx_gelu_fusion.cc index 7824e970d2..2b627543eb 100644 --- a/mindspore/lite/tools/optimizer/fusion/onnx_gelu_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/onnx_gelu_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/onnx_gelu_fusion.h" #include "nnacl/op_base.h" diff --git a/mindspore/lite/tools/optimizer/fusion/reshape_reshape_fusion.cc b/mindspore/lite/tools/optimizer/fusion/reshape_reshape_fusion.cc index 94c0b868b2..b65142b7cd 100644 --- a/mindspore/lite/tools/optimizer/fusion/reshape_reshape_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/reshape_reshape_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/reshape_reshape_fusion.h" #include "ops/op_utils.h" #include "ops/reshape.h" @@ -50,7 +51,9 @@ const AnfNodePtr ReshapeReshapeFusion::Process(const FuncGraphPtr &func_graph, c MS_LOG(ERROR) << "Build reshape primitive failed."; return nullptr; } - auto value_node = NewValueNode(reshape_prim); + auto reshape_prim_c = reshape_prim->GetPrim(); + MS_CHECK_TRUE_RET(reshape_prim_c != nullptr, nullptr); + auto value_node = NewValueNode(reshape_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto input = utils::cast((*equiv)[reshape_input_]); auto shape = utils::cast((*equiv)[reshape_shape_]); diff --git a/mindspore/lite/tools/optimizer/fusion/scale_activation_fusion.cc b/mindspore/lite/tools/optimizer/fusion/scale_activation_fusion.cc index 9652833f7f..de80929792 100644 --- a/mindspore/lite/tools/optimizer/fusion/scale_activation_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/scale_activation_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/scale_activation_fusion.h" #include #include "ops/fusion/activation.h" @@ -43,8 +44,10 @@ const AnfNodePtr ScaleActivationFusion::Process(const FuncGraphPtr &func_graph, return nullptr; } MS_CHECK_TRUE_RET(act_node->size() == kInputSizeTwo, nullptr); - auto act_prim = GetValueNode>(act_node->input(FIRST_INPUT)); - MS_CHECK_TRUE_RET(act_prim != nullptr && act_prim->GetAttr(ops::kActivationType) != nullptr, nullptr); + auto act_prim = ops::GetOperator(act_node->input(FIRST_INPUT)); + MS_CHECK_TRUE_RET(act_prim != nullptr, nullptr); + auto act_prim_c = act_prim->GetPrim(); + MS_CHECK_TRUE_RET(act_prim_c != nullptr && act_prim_c->GetAttr(ops::kActivationType) != nullptr, nullptr); if (act_prim->get_activation_type() != mindspore::RELU && act_prim->get_activation_type() != mindspore::RELU6) { return nullptr; } @@ -56,16 +59,17 @@ const AnfNodePtr ScaleActivationFusion::Process(const FuncGraphPtr &func_graph, if (IsMarkedTrainOp(scale_cnode) || IsMultiOutputTensors(func_graph, scale_cnode)) { return nullptr; } - - auto scale_prim = GetValueNode>(scale_cnode->input(FIRST_INPUT)); + auto scale_prim = ops::GetOperator(scale_cnode->input(FIRST_INPUT)); MS_ASSERT(scale_prim != nullptr); + auto scale_prim_c = scale_prim->GetPrim(); + MS_CHECK_TRUE_RET(scale_prim_c != nullptr, nullptr); ActivationType act_type = act_prim->get_activation_type(); - if (scale_prim->GetAttr(ops::kActivationType) != nullptr && scale_prim->get_activation_type() != NO_ACTIVATION) { + if (scale_prim_c->GetAttr(ops::kActivationType) != nullptr && scale_prim->get_activation_type() != NO_ACTIVATION) { auto scale_act = scale_prim->get_activation_type(); MS_CHECK_TRUE_RET(scale_act == RELU || scale_act == RELU6, nullptr); act_type = scale_act == RELU6 ? RELU6 : act_type; } - scale_prim->AddAttr(ops::kActivationType, MakeValue(act_type)); + scale_prim_c->AddAttr(ops::kActivationType, MakeValue(act_type)); return scale_node; } } // namespace mindspore::opt diff --git a/mindspore/lite/tools/optimizer/fusion/scale_base_fusion.cc b/mindspore/lite/tools/optimizer/fusion/scale_base_fusion.cc index f13e223ec1..03004ea503 100644 --- a/mindspore/lite/tools/optimizer/fusion/scale_base_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/scale_base_fusion.cc @@ -14,12 +14,14 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/scale_base_fusion.h" #include #include "tools/common/tensor_util.h" #include "tools/optimizer/common/gllo_utils.h" #include "tools/converter/quant_param_holder.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { int ScaleBaseFusion::CalNewCnodeScale(const CNodePtr &curr_cnode, @@ -116,7 +118,7 @@ bool ScaleBaseFusion::CheckCurrCnodeProper(const CNodePtr &scale_cnode) const { return false; } - auto scale_prim = GetValueNode>(scale_cnode->input(0)); + auto scale_prim = ops::GetOperator(scale_cnode->input(0)); MS_CHECK_TRUE_RET(scale_prim != nullptr, false); auto axis_attr = scale_prim->GetAttr(ops::kAxis); MS_CHECK_TRUE_RET(axis_attr != nullptr, false); diff --git a/mindspore/lite/tools/optimizer/fusion/scale_scale_fusion.cc b/mindspore/lite/tools/optimizer/fusion/scale_scale_fusion.cc index 44950f29d0..2c87c0a042 100644 --- a/mindspore/lite/tools/optimizer/fusion/scale_scale_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/scale_scale_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/scale_scale_fusion.h" #include #include @@ -23,6 +24,7 @@ #include "ops/fusion/scale_fusion.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -50,9 +52,11 @@ bool ScaleScaleFusion::CheckScaleNode(const CNodePtr &scale_cnode) const { return false; } MS_CHECK_TRUE_RET(scale_cnode->size() >= kScaleNoBiasLen, false); - auto scale_prim = GetValueNode>(scale_cnode->input(FIRST_INPUT)); + auto scale_prim = ops::GetOperator(scale_cnode->input(FIRST_INPUT)); MS_CHECK_TRUE_RET(scale_prim != nullptr, false); - auto quant_attr = scale_prim->GetAttr("quant_params"); + auto scale_prim_c = scale_prim->GetPrim(); + MS_CHECK_TRUE_RET(scale_prim_c != nullptr, false); + auto quant_attr = scale_prim_c->GetAttr("quant_params"); if (quant_attr != nullptr) { auto quant_param_holder = quant_attr->cast(); MS_CHECK_TRUE_RET(quant_param_holder != nullptr, false); @@ -92,12 +96,16 @@ int ScaleScaleFusion::GetInputParamsAndTensors(const CNodePtr &up_scale_cnode, c } MS_CHECK_TRUE_RET(!scale_input_shape_.empty(), lite::RET_ERROR); - auto up_scale_prim = GetValueNode>(up_scale_cnode->input(FIRST_INPUT)); - MS_CHECK_TRUE_RET(up_scale_prim != nullptr && up_scale_prim->GetAttr(ops::kAxis), lite::RET_ERROR); + auto up_scale_prim = ops::GetOperator(up_scale_cnode->input(FIRST_INPUT)); + MS_CHECK_TRUE_RET(up_scale_prim != nullptr, lite::RET_ERROR); + auto up_scale_prim_c = up_scale_prim->GetPrim(); + MS_CHECK_TRUE_RET(up_scale_prim_c != nullptr && up_scale_prim_c->GetAttr(ops::kAxis), lite::RET_ERROR); auto axis = up_scale_prim->get_axis(); up_scale_axis_ = axis < 0 ? axis + scale_input_shape_.size() : axis; - auto down_scale_prim = GetValueNode>(down_scale_cnode->input(FIRST_INPUT)); - MS_CHECK_TRUE_RET(down_scale_prim != nullptr && down_scale_prim->GetAttr(ops::kAxis), lite::RET_ERROR); + auto down_scale_prim = ops::GetOperator(down_scale_cnode->input(FIRST_INPUT)); + MS_CHECK_TRUE_RET(down_scale_prim != nullptr, lite::RET_ERROR); + auto down_scale_prim_c = down_scale_prim->GetPrim(); + MS_CHECK_TRUE_RET(down_scale_prim_c != nullptr && down_scale_prim_c->GetAttr(ops::kAxis), lite::RET_ERROR); axis = down_scale_prim->get_axis(); down_scale_axis_ = axis < 0 ? axis + scale_input_shape_.size() : axis; @@ -282,9 +290,11 @@ const AnfNodePtr ScaleScaleFusion::Process(const FuncGraphPtr &func_graph, const if (!CheckScaleNode(up_scale_cnode) || !CheckScaleNode(down_scale_cnode)) { return nullptr; } - auto scale_prim = GetValueNode>(up_scale_cnode->input(FIRST_INPUT)); + auto scale_prim = ops::GetOperator(up_scale_cnode->input(FIRST_INPUT)); MS_CHECK_TRUE_RET(scale_prim != nullptr, nullptr); - if (scale_prim->GetAttr(ops::kActivationType) != nullptr && scale_prim->get_activation_type() != NO_ACTIVATION) { + auto scale_prim_c = scale_prim->GetPrim(); + MS_CHECK_TRUE_RET(scale_prim_c != nullptr, nullptr); + if (scale_prim_c->GetAttr(ops::kActivationType) != nullptr && scale_prim->get_activation_type() != NO_ACTIVATION) { return nullptr; } @@ -297,8 +307,10 @@ const AnfNodePtr ScaleScaleFusion::Process(const FuncGraphPtr &func_graph, const MS_LOG(ERROR) << "Generate new weight parameter node failed."; return nullptr; } - auto down_scale_prim = GetValueNode>(down_scale_cnode->input(FIRST_INPUT)); - MS_CHECK_TRUE_RET(down_scale_prim != nullptr && down_scale_prim->GetAttr(ops::kAxis) != nullptr, nullptr); + auto down_scale_prim = ops::GetOperator(down_scale_cnode->input(FIRST_INPUT)); + MS_CHECK_TRUE_RET(down_scale_prim != nullptr, nullptr); + auto down_scale_prim_c = down_scale_prim->GetPrim(); + MS_CHECK_TRUE_RET(down_scale_prim_c != nullptr && down_scale_prim_c->GetAttr(ops::kAxis) != nullptr, nullptr); down_scale_prim->set_axis(MSMIN(up_scale_axis_, down_scale_axis_)); auto manager = func_graph->manager(); diff --git a/mindspore/lite/tools/optimizer/fusion/sigmoid_mul_fusion.cc b/mindspore/lite/tools/optimizer/fusion/sigmoid_mul_fusion.cc index b2f90a629c..a447a5283f 100644 --- a/mindspore/lite/tools/optimizer/fusion/sigmoid_mul_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/sigmoid_mul_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/sigmoid_mul_fusion.h" #include #include "ops/fusion/activation.h" @@ -50,7 +52,7 @@ const AnfNodePtr SigmoidMulFusion::Process(const FuncGraphPtr &func_graph, const return nullptr; } // activation must sigmoid - auto activation_prim = GetValueNode>(activation_cnode->input(0)); + auto activation_prim = ops::GetOperator(activation_cnode->input(0)); MS_CHECK_TRUE_RET(activation_prim != nullptr, nullptr); if (activation_prim == nullptr || (activation_prim->GetAttr(ops::kActivationType) != nullptr && activation_prim->get_activation_type() != mindspore::SIGMOID)) { diff --git a/mindspore/lite/tools/optimizer/fusion/squeeze_fusion.cc b/mindspore/lite/tools/optimizer/fusion/squeeze_fusion.cc index 27cef39e03..9a788ccc2f 100644 --- a/mindspore/lite/tools/optimizer/fusion/squeeze_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/squeeze_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/squeeze_fusion.h" #include #include "schema/inner/model_generated.h" @@ -21,6 +22,7 @@ #include "ops/unsqueeze.h" #include "tools/optimizer/common/gllo_utils.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { const BaseRef SqueezeFusion::DefinePattern() const { @@ -91,8 +93,8 @@ const AnfNodePtr SqueezeFusion::Process(const FuncGraphPtr &func_graph, const An MS_LOG(ERROR) << "The squeeze or unsqueeze node has no axis value."; return nullptr; } - auto unsqueeze_prim = utils::cast>(unsqueeze_primitive); - auto squeeze_prim = utils::cast>(squeeze_primitive); + auto unsqueeze_prim = api::MakeShared(unsqueeze_primitive); + auto squeeze_prim = api::MakeShared(squeeze_primitive); MS_ASSERT(unsqueeze_prim != nullptr); MS_ASSERT(squeeze_prim != nullptr); if (squeeze_prim->get_axis() == unsqueeze_prim->get_axis()) { diff --git a/mindspore/lite/tools/optimizer/fusion/tensor_dot_fusion.cc b/mindspore/lite/tools/optimizer/fusion/tensor_dot_fusion.cc index 1a30cad461..ece1a40390 100644 --- a/mindspore/lite/tools/optimizer/fusion/tensor_dot_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/tensor_dot_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/tensor_dot_fusion.h" #include #include diff --git a/mindspore/lite/tools/optimizer/fusion/tf_bidirection_gru_fusion.cc b/mindspore/lite/tools/optimizer/fusion/tf_bidirection_gru_fusion.cc index 74ef478532..e1fb5c8c8d 100644 --- a/mindspore/lite/tools/optimizer/fusion/tf_bidirection_gru_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/tf_bidirection_gru_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/tf_bidirection_gru_fusion.h" #include #include @@ -27,6 +29,7 @@ #include "include/common/utils/utils.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -566,8 +569,10 @@ CNodePtr TfBidirectionGruFusion::GetStackedHiddenState(const FuncGraphPtr &func_ MS_ASSERT(bw_init_state != nullptr); auto stack_prim = std::make_shared(); MS_CHECK_TRUE_RET(stack_prim != nullptr, nullptr); + auto stack_prim_c = stack_prim->GetPrim(); + MS_CHECK_TRUE_RET(stack_prim_c != nullptr, nullptr); stack_prim->set_axis(0); - auto value_node = NewValueNode(stack_prim); + auto value_node = NewValueNode(stack_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); std::vector new_node_inputs = {value_node, fw_init_state, bw_init_state}; auto new_node = func_graph->NewCNode(new_node_inputs); @@ -587,8 +592,10 @@ CNodePtr TfBidirectionGruFusion::CreateBiDirectionGruNode(const FuncGraphPtr &fu MS_ASSERT(equiv != nullptr); auto gru_prim = std::make_shared(); MS_CHECK_TRUE_RET(gru_prim != nullptr, nullptr); + auto gru_prim_c = gru_prim->GetPrim(); + MS_CHECK_TRUE_RET(gru_prim_c != nullptr, nullptr); gru_prim->set_bidirectional(true); - auto value_node = NewValueNode(gru_prim); + auto value_node = NewValueNode(gru_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto fw_gate_kernel = utils::cast((*equiv)[fw_vars_[var_offset]]); @@ -686,9 +693,11 @@ CNodePtr TfBidirectionGruFusion::GetPostProcessNode(const FuncGraphPtr &func_gra MS_ASSERT(gru_output != nullptr); auto split_prim = std::make_shared(); MS_CHECK_TRUE_RET(split_prim != nullptr, nullptr); + auto split_prim_c = split_prim->GetPrim(); + MS_CHECK_TRUE_RET(split_prim_c != nullptr, nullptr); split_prim->set_output_num(2); split_prim->set_axis(1); - auto split_value_node = NewValueNode(split_prim); + auto split_value_node = NewValueNode(split_prim_c); MS_CHECK_TRUE_RET(split_value_node != nullptr, nullptr); std::vector new_node_inputs = {split_value_node, gru_output}; auto split_new_node = func_graph->NewCNode(new_node_inputs); @@ -709,8 +718,10 @@ CNodePtr TfBidirectionGruFusion::GetPostProcessNode(const FuncGraphPtr &func_gra auto concat_prim = std::make_shared(); MS_CHECK_TRUE_RET(concat_prim != nullptr, nullptr); + auto concat_prim_c = concat_prim->GetPrim(); + MS_CHECK_TRUE_RET(concat_prim_c != nullptr, nullptr); concat_prim->set_axis(3); - auto concat_value_node = NewValueNode(concat_prim); + auto concat_value_node = NewValueNode(concat_prim_c); MS_CHECK_TRUE_RET(concat_value_node != nullptr, nullptr); std::vector concat_new_node_inputs = {concat_value_node, split_out1, split_out2}; auto concat_new_node = func_graph->NewCNode(concat_new_node_inputs); @@ -722,8 +733,10 @@ CNodePtr TfBidirectionGruFusion::GetPostProcessNode(const FuncGraphPtr &func_gra auto squeeze_prim = std::make_shared(); MS_CHECK_TRUE_RET(squeeze_prim != nullptr, nullptr); + auto squeeze_prim_c = squeeze_prim->GetPrim(); + MS_CHECK_TRUE_RET(squeeze_prim_c != nullptr, nullptr); squeeze_prim->set_axis(std::vector{1}); - auto squeeze_value_node = NewValueNode(squeeze_prim); + auto squeeze_value_node = NewValueNode(squeeze_prim_c); MS_CHECK_TRUE_RET(squeeze_value_node != nullptr, nullptr); std::vector squeeze_new_node_inputs = {squeeze_value_node, concat_new_node}; auto squeeze_new_node = func_graph->NewCNode(squeeze_new_node_inputs); @@ -737,7 +750,9 @@ CNodePtr TfBidirectionGruFusion::GetPostProcessNode(const FuncGraphPtr &func_gra MS_CHECK_TRUE_RET(transpose_prim != nullptr, nullptr); auto transpose_perm = BuildIntVecParameterNode(func_graph, {1, 0, 2}, "transpose_" + base_name + "_perm"); MS_CHECK_TRUE_RET(transpose_perm != nullptr, nullptr); - auto transpose_new_node = func_graph->NewCNode(transpose_prim, {squeeze_new_node, transpose_perm}); + auto transpose_prim_c = transpose_prim->GetPrim(); + MS_CHECK_TRUE_RET(transpose_prim_c != nullptr, nullptr); + auto transpose_new_node = func_graph->NewCNode(transpose_prim_c, {squeeze_new_node, transpose_perm}); MS_CHECK_TRUE_RET(transpose_new_node != nullptr, nullptr); transpose_new_node->set_fullname_with_scope("transpose_" + base_name); if (gru_output->abstract() != nullptr) { diff --git a/mindspore/lite/tools/optimizer/fusion/tf_gelu_fusion.cc b/mindspore/lite/tools/optimizer/fusion/tf_gelu_fusion.cc index 3efcd0be65..926ff34f81 100644 --- a/mindspore/lite/tools/optimizer/fusion/tf_gelu_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/tf_gelu_fusion.cc @@ -14,9 +14,11 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/tf_gelu_fusion.h" #include "ops/op_utils.h" #include "nnacl/op_base.h" +#include "mindapi/base/types.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/fusion/tf_lstm_cell_fusion.cc b/mindspore/lite/tools/optimizer/fusion/tf_lstm_cell_fusion.cc index aebfd1e4ed..61192be612 100644 --- a/mindspore/lite/tools/optimizer/fusion/tf_lstm_cell_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/tf_lstm_cell_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/tf_lstm_cell_fusion.h" #include #include "ops/lstm.h" @@ -357,10 +359,12 @@ CNodePtr TfLstmCellFusion::CreateLSTMNode(const FuncGraphPtr &func_graph, const MS_ASSERT(equiv != nullptr); auto lstm_prim = std::make_shared(); MS_CHECK_TRUE_RET(lstm_prim != nullptr, nullptr); + auto lstm_prim_c = lstm_prim->GetPrim(); + MS_CHECK_TRUE_RET(lstm_prim_c != nullptr, nullptr); lstm_prim->set_bidirectional(false); lstm_prim->set_zoneout_cell(zoneout_cell); lstm_prim->set_zoneout_hidden(zoneout_hidden); - auto value_node = NewValueNode(lstm_prim); + auto value_node = NewValueNode(lstm_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto &vars = while_input_vars_; diff --git a/mindspore/lite/tools/optimizer/fusion/tflite_lstm_cell_fusion.cc b/mindspore/lite/tools/optimizer/fusion/tflite_lstm_cell_fusion.cc index 1351d38615..ce24a714e1 100644 --- a/mindspore/lite/tools/optimizer/fusion/tflite_lstm_cell_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/tflite_lstm_cell_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/tflite_lstm_cell_fusion.h" #include #include @@ -522,10 +524,12 @@ CNodePtr TfliteLstmCellFusion::CreateLSTMNode(const FuncGraphPtr &func_graph, co */ auto lstm_prim = std::make_shared(); MS_CHECK_TRUE_RET(lstm_prim != nullptr, nullptr); + auto lstm_prim_c = lstm_prim->GetPrim(); + MS_CHECK_TRUE_RET(lstm_prim_c != nullptr, nullptr); lstm_prim->set_bidirectional(false); lstm_prim->set_zoneout_cell(zoneout_cell); lstm_prim->set_zoneout_hidden(zoneout_hidden); - auto value_node = NewValueNode(lstm_prim); + auto value_node = NewValueNode(lstm_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto &vars = while_input_vars_; @@ -614,7 +618,9 @@ CNodePtr TfliteLstmCellFusion::CreateOutputGetItem(const FuncGraphPtr &func_grap MS_LOG(ERROR) << "NewValueNode is nullptr"; return nullptr; } - CNodePtr get_item_cnode = func_graph->NewCNode(tuple_get_item_prim, {node, get_item_value}); + auto tuple_get_item_prim_c = tuple_get_item_prim->GetPrim(); + MS_ASSERT(tuple_get_item_prim_c != nullptr); + CNodePtr get_item_cnode = func_graph->NewCNode(tuple_get_item_prim_c, {node, get_item_value}); MS_CHECK_TRUE_RET(get_item_cnode != nullptr, nullptr); auto abstract = lite::CreateTensorAbstract({}, kNumberTypeFloat32); if (abstract == nullptr) { @@ -714,11 +720,13 @@ CNodePtr TfliteLstmCellFusion::CreateSqueezeNode(const FuncGraphPtr &func_graph, MS_ASSERT(func_graph != nullptr && input_node != nullptr); auto squeeze_prim = std::make_shared(); MS_CHECK_TRUE_RET(squeeze_prim != nullptr, nullptr); + auto squeeze_prim_c = squeeze_prim->GetPrim(); + MS_CHECK_TRUE_RET(squeeze_prim_c != nullptr, nullptr); std::vector axis_vec; std::transform(axis.begin(), axis.end(), std::back_inserter(axis_vec), [](int val) { return static_cast(val); }); squeeze_prim->set_axis(axis_vec); - auto squeeze_cnode = func_graph->NewCNode(squeeze_prim, {input_node}); + auto squeeze_cnode = func_graph->NewCNode(squeeze_prim_c, {input_node}); MS_CHECK_TRUE_RET(squeeze_cnode != nullptr, nullptr); if (input_node->abstract() != nullptr) { squeeze_cnode->set_abstract(input_node->abstract()->Clone()); diff --git a/mindspore/lite/tools/optimizer/fusion/tflite_rel_pos_multi_head_attention_fusion.cc b/mindspore/lite/tools/optimizer/fusion/tflite_rel_pos_multi_head_attention_fusion.cc index b430ee1049..9d8b740c08 100644 --- a/mindspore/lite/tools/optimizer/fusion/tflite_rel_pos_multi_head_attention_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/tflite_rel_pos_multi_head_attention_fusion.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/tflite_rel_pos_multi_head_attention_fusion.h" #include #include @@ -217,11 +219,13 @@ CNodePtr TfliteRelPosMultiHeadAttentionFusion::CreateRelPosMultiHeadAttentionNod MS_ASSERT(func_graph != nullptr && equiv != nullptr); auto attention_prim = BuildAttentionPrim(equiv); MS_CHECK_TRUE_RET(attention_prim != nullptr, nullptr); - if (SetQuantParamForAttentionNode(attention_prim, equiv) != lite::RET_OK) { + auto attention_prim_c = attention_prim->GetPrim(); + MS_CHECK_TRUE_RET(attention_prim_c != nullptr, nullptr); + if (SetQuantParamForAttentionNode(attention_prim_c, equiv) != lite::RET_OK) { MS_LOG(ERROR) << "set quant param for attehtion node failed."; return nullptr; } - auto value_node = NewValueNode(attention_prim); + auto value_node = NewValueNode(attention_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); auto input_q = utils::cast((*equiv)[input_q_]); auto input_k = utils::cast((*equiv)[input_k_]); @@ -231,26 +235,28 @@ CNodePtr TfliteRelPosMultiHeadAttentionFusion::CreateRelPosMultiHeadAttentionNod auto weight_q = utils::cast((*equiv)[weight_q_]); auto transpose_prim = std::make_shared(); MS_CHECK_TRUE_RET(transpose_prim != nullptr, nullptr); + auto transpose_prim_c = transpose_prim->GetPrim(); + MS_CHECK_TRUE_RET(transpose_prim_c != nullptr, nullptr); auto transpose_perm = BuildIntVecParameterNode(func_graph, {1, 0}, "transpose" + base_name + "_perm"); MS_CHECK_TRUE_RET(transpose_perm != nullptr, nullptr); - auto weight_q_transpose = func_graph->NewCNode(transpose_prim, {weight_q, transpose_perm}); + auto weight_q_transpose = func_graph->NewCNode(transpose_prim_c, {weight_q, transpose_perm}); MS_CHECK_TRUE_RET(weight_q_transpose != nullptr, nullptr); weight_q_transpose->set_fullname_with_scope("transpose_wq" + base_name); auto weight_k = utils::cast((*equiv)[weight_k_]); - auto weight_k_transpose = func_graph->NewCNode(transpose_prim, {weight_k, transpose_perm}); + auto weight_k_transpose = func_graph->NewCNode(transpose_prim_c, {weight_k, transpose_perm}); MS_CHECK_TRUE_RET(weight_k_transpose != nullptr, nullptr); weight_k_transpose->set_fullname_with_scope("transpose_wk" + base_name); auto weight_v = utils::cast((*equiv)[weight_v_]); - auto weight_v_transpose = func_graph->NewCNode(transpose_prim, {weight_v, transpose_perm}); + auto weight_v_transpose = func_graph->NewCNode(transpose_prim_c, {weight_v, transpose_perm}); MS_CHECK_TRUE_RET(weight_v_transpose != nullptr, nullptr); weight_v_transpose->set_fullname_with_scope("transpose_wv" + base_name); auto weight_p = utils::cast((*equiv)[weight_p_]); auto weight_o = utils::cast((*equiv)[weight_o_]); - auto weight_o_transpose = func_graph->NewCNode(transpose_prim, {weight_o, transpose_perm}); + auto weight_o_transpose = func_graph->NewCNode(transpose_prim_c, {weight_o, transpose_perm}); MS_CHECK_TRUE_RET(weight_o_transpose != nullptr, nullptr); weight_o_transpose->set_fullname_with_scope("transpose_wo" + base_name); diff --git a/mindspore/lite/tools/optimizer/fusion/transpose_fusion.cc b/mindspore/lite/tools/optimizer/fusion/transpose_fusion.cc index e5d4893d74..9b08cd29d5 100644 --- a/mindspore/lite/tools/optimizer/fusion/transpose_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/transpose_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/transpose_fusion.h" #include #include @@ -23,6 +24,7 @@ #include "tools/optimizer/common/format_utils.h" #include "ops/fusion/scale_fusion.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { bool IsBNCNode(const BaseRef &n) { @@ -160,7 +162,9 @@ CNodePtr GenTransposeNode(const FuncGraphPtr &func_graph, const AnfNodePtr &inpu MS_ASSERT(func_graph != nullptr && input_node != nullptr); auto trans_prim = std::make_shared(); MS_CHECK_TRUE_RET(trans_prim != nullptr, nullptr); - auto cnode = func_graph->NewCNode(trans_prim, {input_node, perm}); + auto trans_prim_c = trans_prim->GetPrim(); + MS_CHECK_TRUE_RET(trans_prim_c != nullptr, nullptr); + auto cnode = func_graph->NewCNode(trans_prim_c, {input_node, perm}); MS_CHECK_TRUE_RET(cnode != nullptr, nullptr); cnode->set_fullname_with_scope(cnode_name); auto quant_params_holder = std::make_shared(2, 1); diff --git a/mindspore/lite/tools/optimizer/fusion/transpose_matmul_fusion.cc b/mindspore/lite/tools/optimizer/fusion/transpose_matmul_fusion.cc index 927e896267..a8c01158d9 100644 --- a/mindspore/lite/tools/optimizer/fusion/transpose_matmul_fusion.cc +++ b/mindspore/lite/tools/optimizer/fusion/transpose_matmul_fusion.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/fusion/transpose_matmul_fusion.h" #include #include @@ -83,7 +84,7 @@ bool TransposeMatMulFusion::Run(const FuncGraphPtr &func_graph) { if (!CheckInputTransposeNode(func_graph, cnode, indices_need_fuse, sizeof(bool) * DIMENSION_2D)) { continue; } - auto matmul_prim = GetValueNode>(cnode->input(0)); + auto matmul_prim = ops::GetOperator(cnode->input(0)); MS_ASSERT(matmul_prim != nullptr); auto manager = func_graph->manager(); MS_ASSERT(manager != nullptr); diff --git a/mindspore/lite/tools/optimizer/graph/add_tensor_array.cc b/mindspore/lite/tools/optimizer/graph/add_tensor_array.cc index 43972b92a5..5582576e81 100644 --- a/mindspore/lite/tools/optimizer/graph/add_tensor_array.cc +++ b/mindspore/lite/tools/optimizer/graph/add_tensor_array.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/add_tensor_array.h" #include #include @@ -84,7 +85,9 @@ static int SetGraphOutput(const FuncGraphPtr &func_graph, const AnfNodePtr &tens MS_LOG(ERROR) << "make_tuple_prim_ptr is nullptr"; return lite::RET_NULL_PTR; } - auto make_tuple_vnode = NewValueNode(make_tuple_prim_ptr); + auto make_tuple_prim_c = make_tuple_prim_ptr->GetPrim(); + MS_CHECK_TRUE_RET(make_tuple_prim_c != nullptr, lite::RET_NULL_PTR); + auto make_tuple_vnode = NewValueNode(make_tuple_prim_c); MS_CHECK_TRUE_RET(make_tuple_vnode != nullptr, lite::RET_NULL_PTR); auto make_tuple_cnode = func_graph->NewCNode({make_tuple_vnode, output_node, tensor_array_write_node}); if (make_tuple_cnode == nullptr) { @@ -99,7 +102,9 @@ static int SetGraphOutput(const FuncGraphPtr &func_graph, const AnfNodePtr &tens MS_LOG(ERROR) << "return_prim_ptr is nullptr"; return lite::RET_NULL_PTR; } - auto return_value_node = NewValueNode(return_prim_ptr); + auto return_prim_c = return_prim_ptr->GetPrim(); + MS_CHECK_TRUE_RET(return_prim_c != nullptr, lite::RET_NULL_PTR); + auto return_value_node = NewValueNode(return_prim_c); MS_CHECK_TRUE_RET(return_value_node != nullptr, lite::RET_NULL_PTR); auto new_return_node = func_graph->NewCNode({return_value_node, make_tuple_cnode}); MS_CHECK_TRUE_RET(new_return_node != nullptr, lite::RET_NULL_PTR); @@ -165,12 +170,14 @@ const AnfNodePtr AddTensorArray::Process(const FuncGraphPtr &func_graph, const A // tensor_array auto tensor_array = std::make_shared(); MS_CHECK_TRUE_RET(tensor_array != nullptr, nullptr); + auto tensor_array_c = tensor_array->GetPrim(); + MS_CHECK_TRUE_RET(tensor_array_c != nullptr, nullptr); std::vector element_shape; std::for_each(tensor_info->shape().begin(), tensor_info->shape().end(), [&element_shape](int64_t v) { element_shape.push_back(static_cast(v)); }); tensor_array->set_element_shape(element_shape); tensor_array->set_data_type(tensor_info->data_type()); - auto tensor_array_vnode = NewValueNode(tensor_array); + auto tensor_array_vnode = NewValueNode(tensor_array_c); MS_CHECK_TRUE_RET(tensor_array_vnode != nullptr, nullptr); auto num_tensors_vnode = NewValueNode(kDefaultNumTensors); MS_CHECK_TRUE_RET(num_tensors_vnode != nullptr, nullptr); @@ -182,7 +189,9 @@ const AnfNodePtr AddTensorArray::Process(const FuncGraphPtr &func_graph, const A // {"handle", "index", "flow_in"} -> {"tensor"} auto tensor_array_read = std::make_shared(); MS_CHECK_TRUE_RET(tensor_array_read != nullptr, nullptr); - auto tensor_array_read_vnode = NewValueNode(tensor_array_read); + auto tensor_array_read_c = tensor_array_read->GetPrim(); + MS_CHECK_TRUE_RET(tensor_array_read_c != nullptr, nullptr); + auto tensor_array_read_vnode = NewValueNode(tensor_array_read_c); MS_CHECK_TRUE_RET(tensor_array_read_vnode != nullptr, nullptr); auto read_index_vnode = NewValueNode(kDefaultIndex); MS_CHECK_TRUE_RET(read_index_vnode != nullptr, nullptr); @@ -198,7 +207,9 @@ const AnfNodePtr AddTensorArray::Process(const FuncGraphPtr &func_graph, const A // {"handle", "index", "value", "flow_in"} -> {"flow_out"} auto tensor_array_write = std::make_shared(); MS_CHECK_TRUE_RET(tensor_array_write != nullptr, nullptr); - auto tensor_array_write_vnode = NewValueNode(tensor_array_write); + auto tensor_array_write_c = tensor_array_write->GetPrim(); + MS_CHECK_TRUE_RET(tensor_array_write_c != nullptr, nullptr); + auto tensor_array_write_vnode = NewValueNode(tensor_array_write_c); MS_CHECK_TRUE_RET(tensor_array_write_vnode != nullptr, nullptr); auto write_index_vnode = NewValueNode(kDefaultIndex); MS_CHECK_TRUE_RET(write_index_vnode != nullptr, nullptr); diff --git a/mindspore/lite/tools/optimizer/graph/clip_convert_activation_pass.cc b/mindspore/lite/tools/optimizer/graph/clip_convert_activation_pass.cc index bc5ae03934..69cfc315c1 100644 --- a/mindspore/lite/tools/optimizer/graph/clip_convert_activation_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/clip_convert_activation_pass.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/graph/clip_convert_activation_pass.h" #include #include @@ -44,7 +46,7 @@ bool ClipConvertActivationPass::Run(const FuncGraphPtr &graph) { auto clip_cnode = node->cast(); MS_ASSERT(clip_cnode != nullptr); MS_ASSERT(clip_cnode->size() >= kClipMinIndex); - auto clip_c = GetValueNode(clip_cnode->input(0)); + auto clip_c = ops::GetOperator(clip_cnode->input(0)); MS_ASSERT(clip_c != nullptr); float max = FLT_MAX; float min = -FLT_MAX; @@ -88,7 +90,9 @@ bool ClipConvertActivationPass::Run(const FuncGraphPtr &graph) { if (min == 0 && max == FLT_MAX) { primitive_c->set_activation_type(mindspore::RELU); } - auto value_node = NewValueNode(primitive_c); + auto primitive = primitive_c->GetPrim(); + MS_CHECK_TRUE_MSG(primitive != nullptr, false, "primitive is nullptr"); + auto value_node = NewValueNode(primitive); MS_CHECK_TRUE_MSG(value_node != nullptr, false, "value_node is nullptr"); std::vector op_inputs = {value_node}; op_inputs.push_back(clip_cnode->input(1)); diff --git a/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc b/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc index 4b7c990cbf..1a29a7dfa7 100644 --- a/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/control_flow_pass.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/graph/control_flow_pass.h" #include #include diff --git a/mindspore/lite/tools/optimizer/graph/decrease_transpose_algo.cc b/mindspore/lite/tools/optimizer/graph/decrease_transpose_algo.cc index 9b67e5075b..d6e00600ce 100644 --- a/mindspore/lite/tools/optimizer/graph/decrease_transpose_algo.cc +++ b/mindspore/lite/tools/optimizer/graph/decrease_transpose_algo.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/decrease_transpose_algo.h" #include #include diff --git a/mindspore/lite/tools/optimizer/graph/dump_graph.h b/mindspore/lite/tools/optimizer/graph/dump_graph.h index 6833b0d978..c965708740 100644 --- a/mindspore/lite/tools/optimizer/graph/dump_graph.h +++ b/mindspore/lite/tools/optimizer/graph/dump_graph.h @@ -19,6 +19,7 @@ #include "backend/common/optimizer/pass.h" #include "tools/converter/export_model.h" #include "include/registry/pass_base.h" +#include "mindapi/ir/func_graph.h" namespace mindspore { namespace opt { @@ -37,7 +38,8 @@ class DumpGraph : public registry::PassBase, public Pass { bool Execute(const api::FuncGraphPtr &func_graph) override { MS_CHECK_TRUE_MSG(func_graph != nullptr, false, "funcGraph is a nullptr."); - auto graph = std::dynamic_pointer_cast(func_graph); + auto impl = func_graph->impl(); + auto graph = std::dynamic_pointer_cast(impl); MS_CHECK_TRUE_MSG(graph != nullptr, false, "Graph is a nullptr."); return Run(graph); } diff --git a/mindspore/lite/tools/optimizer/graph/group_depthwise_op_convert_pass.cc b/mindspore/lite/tools/optimizer/graph/group_depthwise_op_convert_pass.cc index fff7da32ea..b7b6ff728d 100644 --- a/mindspore/lite/tools/optimizer/graph/group_depthwise_op_convert_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/group_depthwise_op_convert_pass.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/graph/group_depthwise_op_convert_pass.h" #include #include @@ -24,6 +26,7 @@ #include "tools/common/tensor_util.h" #include "securec/include/securec.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { @@ -47,7 +50,7 @@ bool GroupDepthwiseOpConvertPass::Run(const FuncGraphPtr &graph) { MS_ASSERT(prim_node != nullptr); auto prim_value_node = prim_node->cast(); MS_ASSERT(prim_value_node != nullptr && prim_value_node->value != nullptr); - auto conv2d_fusion = prim_value_node->value()->cast>(); + auto conv2d_fusion = ops::GetOperator(prim_value_node); if (conv2d_fusion == nullptr) { MS_LOG(ERROR) << "the input of depthwiseConv2d is null"; return false; @@ -102,7 +105,7 @@ bool GroupDepthwiseOpConvertPass::Run(const FuncGraphPtr &graph) { status = TransFilterFormat(weight_value, weight_src_format, weight_dst_format); if (status == RET_OK) { - conv2d_fusion->AddAttr(ops::kFormat, MakeValue(weight_dst_format)); + conv2d_fusion->AddAttr(ops::kFormat, api::MakeValue(weight_dst_format)); } else { MS_LOG(ERROR) << "TransFilter " << EnumNameFormat(schema::EnumValuesFormat()[weight_dst_format]) << "To" << EnumNameFormat(weight_dst_format) << " failed, node : " << node->fullname_with_scope(); diff --git a/mindspore/lite/tools/optimizer/graph/infershape_pass.cc b/mindspore/lite/tools/optimizer/graph/infershape_pass.cc index d05daf2a27..4930d0c728 100644 --- a/mindspore/lite/tools/optimizer/graph/infershape_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/infershape_pass.cc @@ -14,11 +14,13 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/infershape_pass.h" #include "tools/common/node_util.h" #include "tools/common/tensor_util.h" #include "nnacl/op_base.h" #include "src/common/log_util.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/graph/lite_tensor_extractor.cc b/mindspore/lite/tools/optimizer/graph/lite_tensor_extractor.cc index 3fbe44d19e..fe98ab902c 100644 --- a/mindspore/lite/tools/optimizer/graph/lite_tensor_extractor.cc +++ b/mindspore/lite/tools/optimizer/graph/lite_tensor_extractor.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/lite_tensor_extractor.h" #include #include diff --git a/mindspore/lite/tools/optimizer/graph/node_infershape.cc b/mindspore/lite/tools/optimizer/graph/node_infershape.cc index c9ad1daeb7..ebbc3602f4 100644 --- a/mindspore/lite/tools/optimizer/graph/node_infershape.cc +++ b/mindspore/lite/tools/optimizer/graph/node_infershape.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/node_infershape.h" #include #include @@ -28,6 +29,7 @@ #include "src/registry/kernel_interface_registry.h" #include "tools/optimizer/graph/lite_tensor_extractor.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/graph/reduce_same_act_pass.cc b/mindspore/lite/tools/optimizer/graph/reduce_same_act_pass.cc index 2ad71a9830..ec2e52fc34 100644 --- a/mindspore/lite/tools/optimizer/graph/reduce_same_act_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/reduce_same_act_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/reduce_same_act_pass.h" #include "ops/op_utils.h" #include "src/common/utils.h" @@ -25,7 +26,7 @@ namespace mindspore { namespace opt { namespace { constexpr size_t kMinUsersSize = 2; -} +} // namespace bool ReduceSameActPass::Run(const FuncGraphPtr &func_graph) { auto node_list = TopoSort(func_graph->get_return()); auto manager = Manage(func_graph, true); @@ -56,7 +57,7 @@ bool ReduceSameActPass::Run(const FuncGraphPtr &func_graph) { if (post_cnode == nullptr) { return false; } - auto primitive_c = GetValueNode>(post_cnode->input(0)); + auto primitive_c = ops::GetOperator(post_cnode->input(0)); if (primitive_c == nullptr) { return false; } diff --git a/mindspore/lite/tools/optimizer/graph/redundant_op_remove_pass.cc b/mindspore/lite/tools/optimizer/graph/redundant_op_remove_pass.cc index b6caae8548..d7e9405915 100644 --- a/mindspore/lite/tools/optimizer/graph/redundant_op_remove_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/redundant_op_remove_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/redundant_op_remove_pass.h" #include #include @@ -105,7 +106,10 @@ int ProcessDependencyWithTwoNodes(const FuncGraphPtr &func_graph, const CNodePtr tr.SetEdge(post_node, iter->second, NewValueNode(std::make_shared())); tr.Commit(); auto depend_prim = std::make_shared(); - auto depend_node = func_graph->NewCNode(depend_prim, {post_node, pre_node}); + MS_ASSERT(depend_prim != nullptr); + auto depend_prim_c = depend_prim->GetPrim(); + MS_ASSERT(depend_prim_c != nullptr); + auto depend_node = func_graph->NewCNode(depend_prim_c, {post_node, pre_node}); MS_CHECK_TRUE_MSG(depend_prim != nullptr, lite::RET_NULL_PTR, "NewCNode Failed"); MS_CHECK_TRUE_MSG(depend_node != nullptr, lite::RET_NULL_PTR, "NewCNode Failed"); depend_node->set_fullname_with_scope(cnode->fullname_with_scope()); @@ -121,7 +125,11 @@ int ProcessInputHaveDependency(const FuncGraphPtr &func_graph, const CNodePtr &c if (ProcessDependencyWithTwoNodes(func_graph, cnode, false) == lite::RET_OK) { return lite::RET_OK; } - auto make_tuple_prim = NewValueNode(std::make_shared()); + auto make_tuple_node = std::make_shared(); + MS_CHECK_TRUE_MSG(make_tuple_node != nullptr, lite::RET_NULL_PTR, "make_tuple_node Failed"); + auto make_tuple_prim_c = make_tuple_node->GetPrim(); + MS_CHECK_TRUE_MSG(make_tuple_prim_c != nullptr, lite::RET_NULL_PTR, "make_tuple_prim_c Failed"); + auto make_tuple_prim = NewValueNode(make_tuple_prim_c); auto manager = func_graph->manager(); MS_CHECK_TRUE_MSG(make_tuple_prim != nullptr, lite::RET_NULL_PTR, "NewCNode Failed"); MS_ASSERT(manager != nullptr); @@ -318,7 +326,7 @@ int RemoveRedundantOpPass::RemoveInvalidPadOp(const AnfNodePtr &anf_node, const is_invalid = false; } } else { - auto pad_prim = utils::cast>(primitive); + auto pad_prim = api::MakeShared(primitive); MS_ASSERT(pad_prim != nullptr); MS_CHECK_TRUE_RET(pad_prim->GetAttr(ops::kPaddings) != nullptr, lite::RET_ERROR); auto pad_data = pad_prim->get_paddings(); diff --git a/mindspore/lite/tools/optimizer/graph/slice_prepose_pass.cc b/mindspore/lite/tools/optimizer/graph/slice_prepose_pass.cc index 0bbdb5d8a8..5ecbcb6d36 100644 --- a/mindspore/lite/tools/optimizer/graph/slice_prepose_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/slice_prepose_pass.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/graph/slice_prepose_pass.h" #include #include @@ -116,32 +118,32 @@ bool IsScalarNode(const AnfNodePtr &nodePtr) { return false; } -std::shared_ptr GetSlice(const CNodePtr &cnode) { +api::SharedPtr GetSlice(const CNodePtr &cnode) { if (cnode == nullptr) { return nullptr; } - return GetValueNode>(cnode->input(0)); + return ops::GetOperator(cnode->input(0)); } -std::shared_ptr GetSoftmax(const CNodePtr &cnode) { +api::SharedPtr GetSoftmax(const CNodePtr &cnode) { if (cnode == nullptr) { return nullptr; } - return GetValueNode>(cnode->input(0)); + return ops::GetOperator(cnode->input(0)); } -std::shared_ptr GetReshape(const CNodePtr &cnode) { +api::SharedPtr GetReshape(const CNodePtr &cnode) { if (cnode == nullptr) { return nullptr; } - return GetValueNode>(cnode->input(0)); + return ops::GetOperator(cnode->input(0)); } -std::shared_ptr GetFc(const CNodePtr &cnode) { +api::SharedPtr GetFc(const CNodePtr &cnode) { if (cnode == nullptr) { return nullptr; } - return GetValueNode>(cnode->input(0)); + return ops::GetOperator(cnode->input(0)); } std::vector GetTransposePerm(const CNodePtr &node) { @@ -227,8 +229,10 @@ ValueNodePtr SlicePreposePass::CreateSliceValueNode(const std::vector & MS_ASSERT(slice_cnode != nullptr); auto new_slice = std::make_shared(); MS_CHECK_TRUE_MSG(new_slice != nullptr, nullptr, "new_slice is nullptr"); + auto new_slice_c = new_slice->GetPrim(); + MS_CHECK_TRUE_MSG(new_slice_c != nullptr, nullptr, "new_slice_c is nullptr"); new_slice->set_axes(axes); - ValueNodePtr value_node = NewValueNode(new_slice); + ValueNodePtr value_node = NewValueNode(new_slice_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "NewValueNode Failed"); return value_node; } @@ -236,14 +240,16 @@ ValueNodePtr SlicePreposePass::CreateSliceValueNode(const std::vector & ValueNodePtr SlicePreposePass::CopySliceValueNode(const CNodePtr &slice_cnode) { MS_ASSERT(graph != nullptr); MS_ASSERT(slice_cnode != nullptr); - auto slice_c = GetValueNode>(slice_cnode->input(0)); + auto slice_c = ops::GetOperator(slice_cnode->input(0)); if (slice_c == nullptr) { MS_LOG(ERROR) << "slice node is nullptr"; return nullptr; } - auto new_slice_c = std::make_shared(); + auto new_slice = std::make_shared(); + MS_CHECK_TRUE_MSG(new_slice != nullptr, nullptr, "new_slice_c is nullptr"); + auto new_slice_c = new_slice->GetPrim(); MS_CHECK_TRUE_MSG(new_slice_c != nullptr, nullptr, "new_slice_c is nullptr"); - new_slice_c->set_axes(slice_c->get_axes()); + new_slice->set_axes(new_slice->get_axes()); ValueNodePtr value_node = NewValueNode(new_slice_c); MS_CHECK_TRUE_MSG(value_node != nullptr, nullptr, "NewValueNode Failed"); return value_node; @@ -376,7 +382,9 @@ CNodePtr SlicePreposePass::CreateReshapeCNode(const FuncGraphPtr &graph, const s MS_LOG(ERROR) << "primitive_c is nullptr"; return nullptr; } - ValueNodePtr value_node = NewValueNode(new_reshape); + auto new_reshape_c = new_reshape->GetPrim(); + MS_CHECK_TRUE_MSG(new_reshape_c != nullptr, nullptr, "new_reshape_c is nullptr"); + ValueNodePtr value_node = NewValueNode(new_reshape_c); if (value_node == nullptr) { return nullptr; } diff --git a/mindspore/lite/tools/optimizer/graph/special_node_postprocess.cc b/mindspore/lite/tools/optimizer/graph/special_node_postprocess.cc index 5ab85da37e..9c18ed2fda 100644 --- a/mindspore/lite/tools/optimizer/graph/special_node_postprocess.cc +++ b/mindspore/lite/tools/optimizer/graph/special_node_postprocess.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/special_node_postprocess.h" #include #include @@ -21,6 +22,7 @@ #include "include/errorcode.h" #include "tools/optimizer/common/format_utils.h" #include "nnacl//op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.cc b/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.cc index 119c5fff8b..0e794acce8 100644 --- a/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.cc +++ b/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.cc @@ -14,12 +14,14 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/specify_graph_input_format.h" #include #include #include "tools/optimizer/common/format_utils.h" #include "src/common/log_adapter.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.h b/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.h index bfc982f93d..0e6679f2eb 100644 --- a/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.h +++ b/mindspore/lite/tools/optimizer/graph/specify_graph_input_format.h @@ -19,6 +19,7 @@ #include "backend/common/optimizer/pass.h" #include "tools/optimizer/common/gllo_utils.h" +#include "include/api/format.h" namespace mindspore { namespace opt { diff --git a/mindspore/lite/tools/optimizer/graph/split_one_pass.cc b/mindspore/lite/tools/optimizer/graph/split_one_pass.cc index 7c7ed710db..5a8673444d 100644 --- a/mindspore/lite/tools/optimizer/graph/split_one_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/split_one_pass.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/split_one_pass.h" #include "ops/op_utils.h" #include "src/common/utils.h" @@ -46,7 +47,7 @@ bool SplitOnePass::Run(const FuncGraphPtr &func_graph) { if (cnode == nullptr) { return false; } - auto primitive_c = GetValueNode>(cnode->input(0)); + auto primitive_c = ops::GetOperator(cnode->input(0)); if (primitive_c == nullptr) { return false; } diff --git a/mindspore/lite/tools/optimizer/graph/transpose_strategy.cc b/mindspore/lite/tools/optimizer/graph/transpose_strategy.cc index bf514041d2..65d118f8f8 100644 --- a/mindspore/lite/tools/optimizer/graph/transpose_strategy.cc +++ b/mindspore/lite/tools/optimizer/graph/transpose_strategy.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/graph/transpose_strategy.h" #include #include @@ -156,7 +157,7 @@ STATUS ChangeOpCrop(const FuncGraphPtr &func_graph, const CNodePtr &cnode, Forma MS_LOG(ERROR) << "trans_type is invalid."; return lite::RET_ERROR; } - auto crop_prim = GetValueNode>(cnode->input(0)); + auto crop_prim = ops::GetOperator(cnode->input(0)); if (crop_prim == nullptr) { MS_LOG(ERROR) << "cnode is invalid."; return lite::RET_ERROR; @@ -275,7 +276,7 @@ STATUS ChangeOpSlice(const FuncGraphPtr &func_graph, const CNodePtr &cnode, Form return lite::RET_NOT_SUPPORT; } int element_num = shape.front(); - auto prim = GetValueNode>(cnode->input(0)); + auto prim = ops::GetOperator(cnode->input(0)); MS_CHECK_TRUE_MSG(prim != nullptr, RET_ERROR, "GetValueNode failed"); std::vector axes; if (prim->GetAttr(ops::kAxes) == nullptr || prim->get_axes().empty()) { @@ -409,7 +410,7 @@ bool TransposeStrategy::CanFusionIfInsert(const FuncGraphPtr &func_graph, const auto total_node_count = in_nodes.size() + out_nodes.size(); bool can_insert = trans_count > total_node_count / kHalfDivisor; if (CheckPrimitiveType(cnode, prim::kPrimActivation)) { - auto prim_act = GetValueNode>(cnode->input(0)); + auto prim_act = ops::GetOperator(cnode->input(0)); MS_CHECK_TRUE_MSG(prim_act != nullptr, false, "GetValueNode Failed"); if (prim_act->get_activation_type() == mindspore::ActivationType::LEAKY_RELU) { can_insert = trans_count >= total_node_count / kHalfDivisor; diff --git a/mindspore/lite/tools/optimizer/graph/unused_transpose_node_remove_pass.cc b/mindspore/lite/tools/optimizer/graph/unused_transpose_node_remove_pass.cc index e54d6ab5fd..a8887313c3 100644 --- a/mindspore/lite/tools/optimizer/graph/unused_transpose_node_remove_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/unused_transpose_node_remove_pass.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/graph/unused_transpose_node_remove_pass.h" #include #include diff --git a/mindspore/lite/tools/optimizer/graph/update_conv2d_param_pass.cc b/mindspore/lite/tools/optimizer/graph/update_conv2d_param_pass.cc index 671c662d4c..50d963cc24 100644 --- a/mindspore/lite/tools/optimizer/graph/update_conv2d_param_pass.cc +++ b/mindspore/lite/tools/optimizer/graph/update_conv2d_param_pass.cc @@ -13,12 +13,15 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/graph/update_conv2d_param_pass.h" #include #include #include #include "ops/fusion/conv2d_fusion.h" #include "mindspore/lite/include/errorcode.h" +#include "ops/op_utils.h" namespace mindspore::opt { namespace { diff --git a/mindspore/lite/tools/optimizer/parallel/conv2d_info.cc b/mindspore/lite/tools/optimizer/parallel/conv2d_info.cc index bfa488a518..9370847215 100644 --- a/mindspore/lite/tools/optimizer/parallel/conv2d_info.cc +++ b/mindspore/lite/tools/optimizer/parallel/conv2d_info.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/conv2d_info.h" #include #include @@ -29,6 +30,7 @@ #include "tools/optimizer/parallel/spliter.h" #include "tools/optimizer/fisson/fisson_util.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" using mindspore::schema::PrimitiveType_Conv2DFusion; namespace mindspore { @@ -91,7 +93,7 @@ int Conv2DInfo::CheckStrategy(const SplitStrategy &strategy) { } int Conv2DInfo::CheckIfSplit() { - auto conv_prim = GetValueNode>(cnode_->input(kAnfPrimitiveIndex)); + auto conv_prim = ops::GetOperator(cnode_->input(kAnfPrimitiveIndex)); MS_ASSERT(conv_prim != nullptr); auto strides = conv_prim->get_stride(); std::vector weight_shape; @@ -154,11 +156,13 @@ AnfNodePtr Conv2DInfo::CreateOutputsOfSplit(const CNodePtr &orig_node, size_t in auto input_shapes = input_shape_iter->second; auto input_shape = input_shapes.front(); - auto conv_prim = GetValueNode>(cnode_->input(kAnfPrimitiveIndex)); + auto conv_prim = ops::GetOperator(cnode_->input(kAnfPrimitiveIndex)); MS_ASSERT(conv_prim != nullptr); // prim of split auto split_prim = std::make_shared(); MS_CHECK_TRUE_RET(split_prim != nullptr, nullptr); + auto split_prim_c = split_prim->GetPrim(); + MS_CHECK_TRUE_RET(split_prim_c != nullptr, nullptr); std::vector new_splits = splits; if (split_mode_ == SplitH) { split_prim->set_extend_top(std::vector(split_num, 0)); @@ -183,7 +187,7 @@ AnfNodePtr Conv2DInfo::CreateOutputsOfSplit(const CNodePtr &orig_node, size_t in split_prim->set_number_split(split_num); split_prim->set_ratio(new_splits); - auto split_primitive = NewValueNode(split_prim); + auto split_primitive = NewValueNode(split_prim_c); MS_CHECK_TRUE_MSG(split_primitive != nullptr, nullptr, "create SplitWithOverlap return nullptr"); std::vector split_inputs = {split_primitive}; // ori_conv_node must only have one input @@ -258,12 +262,12 @@ int Conv2DInfo::InferParallelCNodes() { } name_ = orig_name; parallel_output_nodes_.clear(); - auto conv_prim = GetValueNode>(cnode_->input(kAnfPrimitiveIndex)); + auto conv_prim = ops::GetOperator(cnode_->input(kAnfPrimitiveIndex)); MS_ASSERT(conv_prim != nullptr); return ConstructOutputCNodes(conv_prim, feature_split_outputs, kernel_split_outputs, bias_split_outputs); } -std::shared_ptr Conv2DInfo::GetNewConvPrimitive(const std::shared_ptr &conv_prim, +std::shared_ptr Conv2DInfo::GetNewConvPrimitive(const api::SharedPtr &conv_prim, size_t dev_index, int cin_sum, int cout_sum) { auto prim = std::make_shared(); MS_CHECK_TRUE_RET(prim != nullptr, nullptr); @@ -319,7 +323,7 @@ std::shared_ptr Conv2DInfo::GetNewConvPrimitive(const std::sh return prim; } -int Conv2DInfo::ConstructOutputCNodes(const std::shared_ptr &conv_prim, +int Conv2DInfo::ConstructOutputCNodes(const api::SharedPtr &conv_prim, const std::vector &feature_split_outputs, const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs) { @@ -343,6 +347,8 @@ int Conv2DInfo::ConstructOutputCNodes(const std::shared_ptr & MS_LOG(ERROR) << "Get new convolution primitive failed."; return RET_ERROR; } + auto prim_c = prim->GetPrim(); + MS_CHECK_TRUE_RET(prim_c != nullptr, RET_ERROR); std::vector conv_inputs; // if split Cout, feature will not be splited if (split_mode_ == SplitCOUT) { @@ -363,7 +369,8 @@ int Conv2DInfo::ConstructOutputCNodes(const std::shared_ptr & conv_inputs.push_back(cnode_->input(kAnfConvBias)); } } - auto conv_cnode = func_graph_->NewCNode(prim, conv_inputs); + + auto conv_cnode = func_graph_->NewCNode(prim_c, conv_inputs); if (conv_cnode == nullptr) { MS_LOG(ERROR) << name_ << " : Failed to create parallel Conv2D node " << i; return lite::RET_ERROR; diff --git a/mindspore/lite/tools/optimizer/parallel/conv2d_info.h b/mindspore/lite/tools/optimizer/parallel/conv2d_info.h index 8dc566e640..43e30c0130 100644 --- a/mindspore/lite/tools/optimizer/parallel/conv2d_info.h +++ b/mindspore/lite/tools/optimizer/parallel/conv2d_info.h @@ -37,7 +37,7 @@ class Conv2DInfo : public OperatorInfo { int CheckStrategy(const SplitStrategy &strategy) override; int InferReplaceOp() override; int InferParallelCNodes() override; - virtual int ConstructOutputCNodes(const std::shared_ptr &conv_prim, + virtual int ConstructOutputCNodes(const api::SharedPtr &conv_prim, const std::vector &feature_split_outputs, const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs); @@ -51,7 +51,7 @@ class Conv2DInfo : public OperatorInfo { private: int CheckConv2DPrimitiveType(); int CheckIfSplit(); - std::shared_ptr GetNewConvPrimitive(const std::shared_ptr &conv_prim, + std::shared_ptr GetNewConvPrimitive(const api::SharedPtr &conv_prim, size_t dev_index, int cin_sum, int cout_sum); }; } // namespace opt diff --git a/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.cc b/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.cc index 1bd99cd633..49d119d6bf 100644 --- a/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.cc +++ b/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/depthwise_conv2d_info.h" #include #include @@ -31,6 +32,7 @@ #include "ops/split_with_overlap.h" #include "src/tensor.h" #include "tools/optimizer/parallel/spliter.h" +#include "ops/op_utils.h" using mindspore::schema::PrimitiveType_Conv2DFusion; namespace mindspore { @@ -205,7 +207,7 @@ bool DepthwiseConv2DInfo::CheckSplitOutputs(const std::vector &featu return true; } -void DepthwiseConv2DInfo::AdJustConvPrim(const std::shared_ptr &conv_prim, +void DepthwiseConv2DInfo::AdJustConvPrim(const api::SharedPtr &conv_prim, int64_t *visited_in_channel, int64_t *visited_out_channel, int64_t *visited_group, int output_conv_index) { MS_ASSERT(conv_prim != nullptr && visited_in_channel != nullptr); @@ -268,15 +270,17 @@ void DepthwiseConv2DInfo::AdJustConvPrim(const std::shared_ptr &conv_prim, +void DepthwiseConv2DInfo::AdJustInputs(const api::SharedPtr &conv_prim, const std::vector &feature_split_outputs, const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs, int output_conv_index) { MS_ASSERT(conv_prim != nullptr); + auto conv_prim_c = conv_prim->GetPrim(); + MS_ASSERT(conv_prim_c != nullptr); std::vector tmp_outputs; std::string conv_cnode_name = cnode_->fullname_with_scope(); bool has_bias = cnode_->size() > kBiasIndex + 1; - std::vector conv_inputs = {NewValueNode(conv_prim)}; + std::vector conv_inputs = {NewValueNode(conv_prim_c)}; if (split_mode_ == SplitN || split_mode_ == SplitH) { conv_inputs.push_back(feature_split_outputs.at(output_conv_index)); conv_inputs.push_back(cnode_->input(kWeightIndex + 1)); @@ -306,7 +310,7 @@ void DepthwiseConv2DInfo::AdJustInputs(const std::shared_ptr parallel_output_nodes_.push_back(tmp_outputs[0]); } -int DepthwiseConv2DInfo::ConstructOutputCNodes(const std::shared_ptr &conv_prim, +int DepthwiseConv2DInfo::ConstructOutputCNodes(const api::SharedPtr &conv_prim, const std::vector &feature_split_outputs, const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs) { @@ -333,7 +337,7 @@ AnfNodePtr DepthwiseConv2DInfo::CreateOutputsOfSplit(const CNodePtr &ori_node, s std::vector *split_outputs, size_t split_dim, size_t split_num, const std::vector &splits) { MS_ASSERT(orig_node != nullptr && split_outputs != nullptr); - auto depth_wise_conv_prim = GetValueNode>(cnode_->input(kAnfPrimitiveIndex)); + auto depth_wise_conv_prim = ops::GetOperator(cnode_->input(kAnfPrimitiveIndex)); MS_ASSERT(depth_wise_conv_prim != nullptr); auto ori_node_name = ori_node->fullname_with_scope(); auto graph_node_input_shapes = Spliter::GetInstance()->graph_node_input_shapes(); @@ -353,6 +357,8 @@ AnfNodePtr DepthwiseConv2DInfo::CreateOutputsOfSplit(const CNodePtr &ori_node, s // prim of split auto split_prim = std::make_shared(); MS_CHECK_TRUE_RET(split_prim != nullptr, nullptr); + auto split_prim_c = split_prim->GetPrim(); + MS_CHECK_TRUE_RET(split_prim_c != nullptr, nullptr); std::vector new_splits = splits; MS_CHECK_TRUE_RET(input_shape.size() > static_cast(split_dim), nullptr); if (split_mode_ == SplitH) { @@ -380,7 +386,7 @@ AnfNodePtr DepthwiseConv2DInfo::CreateOutputsOfSplit(const CNodePtr &ori_node, s std::vector split_inputs; // ori_conv_node must only have one feature input split_inputs.push_back(ori_node->input(input_index + 1)); - auto split_cnode = func_graph_->NewCNode(split_prim, split_inputs); + auto split_cnode = func_graph_->NewCNode(split_prim_c, split_inputs); if (split_cnode == nullptr) { MS_LOG(ERROR) << name_ << " : Failed to create split node."; lite::ReturnCode::GetSingleReturnCode()->UpdateReturnCode(lite::RET_NULL_PTR); @@ -491,7 +497,7 @@ int DepthwiseConv2DInfo::InferParallelCNodes() { } } name_ = input_op_name; - auto depth_wise_conv_prim = GetValueNode>(cnode_->input(kAnfPrimitiveIndex)); + auto depth_wise_conv_prim = ops::GetOperator(cnode_->input(kAnfPrimitiveIndex)); MS_ASSERT(depth_wise_conv_prim != nullptr); return ConstructOutputCNodes(depth_wise_conv_prim, feature_split_outputs, kernel_split_outputs, bias_split_outputs); } diff --git a/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.h b/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.h index d745e27788..dfdcbb9cb5 100644 --- a/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.h +++ b/mindspore/lite/tools/optimizer/parallel/depthwise_conv2d_info.h @@ -34,7 +34,7 @@ class DepthwiseConv2DInfo : public Conv2DInfo { int InferReplaceOp() override; int InferParallelCNodes() override; int CheckStrategy(const SplitStrategy &strategy) override; - int ConstructOutputCNodes(const std::shared_ptr &conv_prim, + int ConstructOutputCNodes(const api::SharedPtr &conv_prim, const std::vector &feature_split_outputs, const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs) override; @@ -48,10 +48,10 @@ class DepthwiseConv2DInfo : public Conv2DInfo { const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs); - void AdJustConvPrim(const std::shared_ptr &conv_prim, int64_t *visited_in_channel, + void AdJustConvPrim(const api::SharedPtr &conv_prim, int64_t *visited_in_channel, int64_t *visited_out_channel, int64_t *visited_group, int output_conv_index); - void AdJustInputs(const std::shared_ptr &conv_prim, + void AdJustInputs(const api::SharedPtr &conv_prim, const std::vector &feature_split_outputs, const std::vector &kernel_split_outputs, const std::vector &bias_split_outputs, int output_conv_index); diff --git a/mindspore/lite/tools/optimizer/parallel/multi_conv_info.cc b/mindspore/lite/tools/optimizer/parallel/multi_conv_info.cc index 99f3b921cb..342e96a4ea 100644 --- a/mindspore/lite/tools/optimizer/parallel/multi_conv_info.cc +++ b/mindspore/lite/tools/optimizer/parallel/multi_conv_info.cc @@ -13,6 +13,8 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/multi_conv_info.h" #include #include @@ -20,6 +22,7 @@ #include "ops/fusion/conv2d_fusion.h" #include "tools/optimizer/parallel/split_strategy.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" using mindspore::converter::FmkType; using mindspore::schema::PrimitiveType_Conv2dTransposeFusion; @@ -65,7 +68,7 @@ bool MultiConvSplit::CheckSplitValid() { for (const auto &conv_node : conv_nodes_) { auto conv_cnode = conv_node->cast(); MS_ASSERT(conv_cnode != nullptr); - auto conv_prim = GetValueNode>(conv_cnode->input(kAnfPrimitiveIndex)); + auto conv_prim = ops::GetOperator(conv_cnode->input(kAnfPrimitiveIndex)); MS_ASSERT(conv_prim != nullptr); MS_CHECK_TRUE_RET(conv_prim->GetAttr(ops::kPadMode) != nullptr, false); if (conv_prim->get_pad_mode() != SAME) { @@ -177,7 +180,7 @@ bool MultiConvSplit::SplitSingleConv(const AnfNodePtr &ori_node, const std::vect MS_ASSERT(ori_node != nullptr && outputs_node != nullptr); auto ori_conv_cnode = ori_node->cast(); MS_ASSERT(ori_conv_cnode != nullptr); - auto ori_attr = GetValueNode>(ori_conv_cnode->input(kAnfPrimitiveIndex)); + auto ori_attr = ops::GetOperator(ori_conv_cnode->input(kAnfPrimitiveIndex)); MS_ASSERT(ori_attr != nullptr); for (int output_conv_index = 0; output_conv_index < static_cast(split_info_.out_num); output_conv_index++) { // Create Conv node attr @@ -194,7 +197,9 @@ bool MultiConvSplit::SplitSingleConv(const AnfNodePtr &ori_node, const std::vect AdJustConvPrim(conv_prim, input_shape, output_conv_index); // node inputs std::vector conv_inputs; - conv_inputs.push_back(NewValueNode(conv_prim)); + auto conv_prim_c = conv_prim->GetPrim(); + MS_ASSERT(conv_prim_c != nullptr); + conv_inputs.push_back(NewValueNode(conv_prim_c)); AdJustInputs(ori_node, inputs_node, output_conv_index, &conv_inputs); // create new conv node if (!CreateNewConvNode(ori_node, conv_inputs, output_conv_index, outputs_node)) { @@ -273,8 +278,8 @@ AnfNodePtr MultiConvSplitH::SplitMultiConv(const AnfNodePtr &node) { return MultiConvNHSplit(node); } -void MultiConvSplitH::AdJustConvPrim(const std::shared_ptr &conv_prim, - const ShapeVector &input_shape, int output_conv_index) { +void MultiConvSplitH::AdJustConvPrim(const api::SharedPtr &conv_prim, const ShapeVector &input_shape, + int output_conv_index) { MS_ASSERT(conv_prim != nullptr); MS_ASSERT(input_shape.size() == kInputSizeFour); int64_t input_h = input_shape.at(kAxisH); diff --git a/mindspore/lite/tools/optimizer/parallel/multi_conv_info.h b/mindspore/lite/tools/optimizer/parallel/multi_conv_info.h index 500b338b43..9ca3a393c1 100644 --- a/mindspore/lite/tools/optimizer/parallel/multi_conv_info.h +++ b/mindspore/lite/tools/optimizer/parallel/multi_conv_info.h @@ -35,7 +35,7 @@ class MultiConvSplit : public MultiNodeSplit { virtual AnfNodePtr SplitMultiConv(const AnfNodePtr &node) = 0; - virtual void AdJustConvPrim(const std::shared_ptr &ori_attr, const ShapeVector &input_shape, + virtual void AdJustConvPrim(const api::SharedPtr &ori_attr, const ShapeVector &input_shape, int output_conv_index) = 0; virtual AnfNodePtr MultiConvNHSplit(const AnfNodePtr &node); @@ -73,7 +73,7 @@ class MultiConvSplitN final : public MultiConvSplit { ~MultiConvSplitN() = default; AnfNodePtr SplitMultiConv(const AnfNodePtr &node) override; - void AdJustConvPrim(const std::shared_ptr &ori_attr, const ShapeVector &input_shape, + void AdJustConvPrim(const api::SharedPtr &ori_attr, const ShapeVector &input_shape, int output_conv_index) override {} }; @@ -84,7 +84,7 @@ class MultiConvSplitCIN final : public MultiConvSplit { ~MultiConvSplitCIN() = default; AnfNodePtr SplitMultiConv(const AnfNodePtr &node) override; - void AdJustConvPrim(const std::shared_ptr &ori_attr, const ShapeVector &input_shape, + void AdJustConvPrim(const api::SharedPtr &ori_attr, const ShapeVector &input_shape, int output_conv_index) override {} }; @@ -96,7 +96,7 @@ class MultiConvSplitCOUT final : public MultiConvSplit { ~MultiConvSplitCOUT() = default; AnfNodePtr SplitMultiConv(const AnfNodePtr &node) override; - void AdJustConvPrim(const std::shared_ptr &ori_attr, const ShapeVector &input_shape, + void AdJustConvPrim(const api::SharedPtr &ori_attr, const ShapeVector &input_shape, int output_conv_index) override {} }; @@ -107,7 +107,7 @@ class MultiConvSplitH final : public MultiConvSplit { ~MultiConvSplitH() = default; AnfNodePtr SplitMultiConv(const AnfNodePtr &node) override; - void AdJustConvPrim(const std::shared_ptr &ori_attr, const ShapeVector &input_shape, + void AdJustConvPrim(const api::SharedPtr &ori_attr, const ShapeVector &input_shape, int output_conv_index) override; }; diff --git a/mindspore/lite/tools/optimizer/parallel/multi_node_split.cc b/mindspore/lite/tools/optimizer/parallel/multi_node_split.cc index fad75a3746..b9e86cdcd4 100644 --- a/mindspore/lite/tools/optimizer/parallel/multi_node_split.cc +++ b/mindspore/lite/tools/optimizer/parallel/multi_node_split.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/multi_node_split.h" #include "tools/optimizer/parallel/multi_conv_info.h" #include "nnacl/op_base.h" diff --git a/mindspore/lite/tools/optimizer/parallel/operator_info.cc b/mindspore/lite/tools/optimizer/parallel/operator_info.cc index 537f575bc2..c561bad6c6 100644 --- a/mindspore/lite/tools/optimizer/parallel/operator_info.cc +++ b/mindspore/lite/tools/optimizer/parallel/operator_info.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/operator_info.h" #include #include "tools/optimizer/parallel/split_strategy.h" @@ -24,6 +25,7 @@ #include "base/core_ops.h" #include "include/errorcode.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { @@ -119,7 +121,9 @@ int OperatorInfo::CreateMultipleOutputsOfAnfNode(const AnfNodePtr &node, size_t auto abstract_scalar = std::make_shared(index); MS_CHECK_TRUE_RET(abstract_scalar != nullptr, lite::RET_ERROR); idx->set_abstract(abstract_scalar); - auto tuple_getitem = func_graph_->NewCNode({NewValueNode(std::make_shared()), node, idx}); + auto tuple_node = std::make_shared(); + auto tuple_prim_c = tuple_node->GetPrim(); + auto tuple_getitem = func_graph_->NewCNode({NewValueNode(tuple_prim_c), node, idx}); if (tuple_getitem == nullptr) { MS_LOG(ERROR) << name_ << " : Failed to create output nodes."; return lite::RET_ERROR; @@ -142,8 +146,10 @@ AnfNodePtr OperatorInfo::CreateConcateNode(const CNodePtr &orig_node, const std: } auto concat_prim = std::make_shared(); MS_CHECK_TRUE_RET(concat_prim != nullptr, nullptr); + auto concat_prim_c = concat_prim->GetPrim(); + MS_CHECK_TRUE_RET(concat_prim_c != nullptr, nullptr); concat_prim->set_axis(concat_dim); - auto value_node = NewValueNode(concat_prim); + auto value_node = NewValueNode(concat_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); std::vector concat_inputs = {value_node}; (void)std::transform(input_nodes.begin(), input_nodes.end(), std::back_inserter(concat_inputs), @@ -170,7 +176,9 @@ AnfNodePtr OperatorInfo::CreateReduceNode(const CNodePtr &orig_node, const std:: // addup inputs element-wise auto addn_prim = std::make_shared(); MS_CHECK_TRUE_RET(addn_prim != nullptr, nullptr); - auto value_node = NewValueNode(addn_prim); + auto addn_prim_c = addn_prim->GetPrim(); + MS_CHECK_TRUE_RET(addn_prim_c != nullptr, nullptr); + auto value_node = NewValueNode(addn_prim_c); MS_CHECK_TRUE_RET(value_node != nullptr, nullptr); std::vector addn_inputs = {value_node}; (void)std::transform(input_nodes.begin(), input_nodes.end(), std::back_inserter(addn_inputs), diff --git a/mindspore/lite/tools/optimizer/parallel/operator_info_register.cc b/mindspore/lite/tools/optimizer/parallel/operator_info_register.cc index d99d46a1a3..6d68889afc 100644 --- a/mindspore/lite/tools/optimizer/parallel/operator_info_register.cc +++ b/mindspore/lite/tools/optimizer/parallel/operator_info_register.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/operator_info_register.h" #include namespace mindspore { diff --git a/mindspore/lite/tools/optimizer/parallel/parallel_pass.cc b/mindspore/lite/tools/optimizer/parallel/parallel_pass.cc index 3a4a71a5cf..12f9260c0a 100644 --- a/mindspore/lite/tools/optimizer/parallel/parallel_pass.cc +++ b/mindspore/lite/tools/optimizer/parallel/parallel_pass.cc @@ -14,19 +14,21 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/parallel_pass.h" #include "include/errorcode.h" #include "ir/tensor.h" #include "tools/optimizer/parallel/operator_info_register.h" #include "ops/fusion/conv2d_fusion.h" #include "nnacl/op_base.h" +#include "ops/op_utils.h" namespace mindspore { namespace opt { namespace { constexpr auto kAnfPrimitiveIndex = 0; -} +} // namespace bool ParallelPass::IsParallelCareNode(const AnfNodePtr &node) { MS_ASSERT(node != nullptr); diff --git a/mindspore/lite/tools/optimizer/parallel/split_strategy.cc b/mindspore/lite/tools/optimizer/parallel/split_strategy.cc index a8ddb923a1..9941e207a2 100644 --- a/mindspore/lite/tools/optimizer/parallel/split_strategy.cc +++ b/mindspore/lite/tools/optimizer/parallel/split_strategy.cc @@ -14,6 +14,7 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/split_strategy.h" #include #include diff --git a/mindspore/lite/tools/optimizer/parallel/spliter.cc b/mindspore/lite/tools/optimizer/parallel/spliter.cc index 4374a6a995..ee8dad20cf 100644 --- a/mindspore/lite/tools/optimizer/parallel/spliter.cc +++ b/mindspore/lite/tools/optimizer/parallel/spliter.cc @@ -14,10 +14,13 @@ * limitations under the License. */ +#define USE_DEPRECATED_API #include "tools/optimizer/parallel/spliter.h" #include #include "tools/optimizer/fisson/fisson_util.h" #include "tools/optimizer/parallel/split_strategy.h" +#include "ops/op_utils.h" + namespace mindspore { namespace opt { Spliter *Spliter::GetInstance() { @@ -57,9 +60,7 @@ void Spliter::VisitNodesOutputs(const FuncGraphPtr &func_graph) { } void Spliter::RecordGraphInfo(const FuncGraphPtr &func_graph) { - if (func_graph == nullptr) { - return; - } + MS_ASSERT(func_graph != nullptr); VisitNodesInputs(func_graph); VisitNodesOutputs(func_graph); for (const auto &node : func_graph->GetOrderedCnodes()) { diff --git a/tests/ut/cpp/pre_activate/common/restore_abs_input_in_backed_infer_test.cc b/tests/ut/cpp/pre_activate/common/restore_abs_input_in_backed_infer_test.cc index 6cab1707b0..a8d0e7335a 100644 --- a/tests/ut/cpp/pre_activate/common/restore_abs_input_in_backed_infer_test.cc +++ b/tests/ut/cpp/pre_activate/common/restore_abs_input_in_backed_infer_test.cc @@ -13,9 +13,12 @@ * See the License for the specific language governing permissions and * limitations under the License. */ + +#define USE_DEPRECATED_API #include #include #include +#include "ops/base_operator.h" #include "ir/primitive.h" #include "include/common/utils/utils.h" #include "abstract/abstract_value.h" @@ -23,15 +26,18 @@ #include "backend/common/optimizer/const_input_to_attr.h" #include "backend/common/optimizer/helper.h" #include "common/common_test.h" + namespace mindspore { namespace opt { -class TestAttr : public ops::PrimitiveC { +class TestAttr : public ops::BaseOperator { public: - TestAttr() : PrimitiveC("") {} + MIND_API_BASE_MEMBER(TestAttr); + TestAttr() : BaseOperator("") {} }; -class TestDynamicInput : public ops::PrimitiveC { +class TestDynamicInput : public ops::BaseOperator { public: - TestDynamicInput() : PrimitiveC("") {} + MIND_API_BASE_MEMBER(TestDynamicInput); + TestDynamicInput() : BaseOperator("") {} }; constexpr auto kAttrConvertTestName = "attr_convert_test"; constexpr auto kDynamicInputTestName = "dynamic_input_test";