From cd2a02199290e1a0a85760fb5240e8d5fcc81aa1 Mon Sep 17 00:00:00 2001 From: yujianfeng Date: Wed, 10 Nov 2021 17:32:59 +0800 Subject: [PATCH] Add shard api --- .../frontend/operator/composite/composite.cc | 71 +++++++++++++++++++ .../frontend/operator/composite/composite.h | 17 +++++ .../operator/ops_front_infer_function.cc | 21 ++++++ .../operator/ops_front_infer_function.h | 2 + .../pipeline/jit/static_analysis/evaluator.cc | 22 ++++++ .../pipeline/jit/static_analysis/evaluator.h | 35 +++++++++ .../jit/static_analysis/static_analysis.cc | 10 +++ .../jit/static_analysis/static_analysis.h | 1 + mindspore/core/abstract/abstract_function.cc | 14 ++++ mindspore/core/abstract/abstract_function.h | 30 ++++++++ mindspore/core/base/core_ops.h | 1 + mindspore/core/utils/trace_info.h | 8 +++ mindspore/ops/composite/__init__.py | 5 +- mindspore/ops/composite/base.py | 44 +++++++++++- mindspore/ops/functional.py | 6 +- tests/ut/cpp/operator/composite_test.cc | 48 +++++++++++++ 16 files changed, 331 insertions(+), 4 deletions(-) diff --git a/mindspore/ccsrc/frontend/operator/composite/composite.cc b/mindspore/ccsrc/frontend/operator/composite/composite.cc index e36b2827bc6..4f9774c8294 100644 --- a/mindspore/ccsrc/frontend/operator/composite/composite.cc +++ b/mindspore/ccsrc/frontend/operator/composite/composite.cc @@ -1115,5 +1115,76 @@ REGISTER_PYBIND_DEFINE(TupleGetItemTensor_, ([](const py::module *m) { *m, "TupleGetItemTensor_") .def(py::init()); })); + +namespace { +FuncGraphPtr GetShard(const AnfNodePtr &shard, const std::vector &origin_graph_params) { + FuncGraphPtr shard_child = std::make_shared(); + shard_child->set_flag(FUNC_GRAPH_FLAG_CORE, true); + + std::vector inputs; + inputs.reserve(origin_graph_params.size() + 1); + (void)inputs.emplace_back(shard); + for (size_t i = 0; i < origin_graph_params.size(); ++i) { + (void)inputs.emplace_back(shard_child->add_parameter()); + } + auto shard_app = shard_child->NewCNodeInOrder(std::move(inputs)); + + shard_child->set_output(shard_app); + return shard_child; +} +} // namespace + +FuncGraphPtr Shard::GenerateFuncGraph(const AbstractBasePtrList &args_spec_list) { + constexpr size_t shard_input_size = 5; + if (args_spec_list.size() != shard_input_size) { + MS_LOG(EXCEPTION) << "'Shard' requires " << shard_input_size + << " inputs. Includes a Cell or function, in_axes, out_axes, device and level."; + } + + MS_EXCEPTION_IF_NULL(args_spec_list[0]); + AbstractFunctionPtr fn = dyn_cast(args_spec_list[0]); + if (fn == nullptr) { + MS_LOG(EXCEPTION) << "'Shard' arg0 must be a 'Function' or 'Cell', but got " << args_spec_list[0]->ToString() + << "."; + } + + auto real_fn = dyn_cast(fn); + MS_EXCEPTION_IF_NULL(real_fn); + FuncGraphPtr origin_graph = real_fn->func_graph(); + MS_EXCEPTION_IF_NULL(origin_graph); + origin_graph->set_flag(FUNC_GRAPH_FLAG_DEFER_INLINE, true); + FuncGraphPtr shard_fg = nullptr; + { + TraceGuard g(std::make_shared(origin_graph->debug_info())); + shard_fg = std::make_shared(); + } + // Create the debug info + auto parameter_size = origin_graph->parameters().size(); + std::ostringstream ss; + ss << "shard{" << parameter_size << "}"; + shard_fg->set_flag(FUNC_GRAPH_FLAG_CORE, true); + shard_fg->debug_info()->set_name(ss.str()); + // Make the Shard node. + std::vector inputs; + inputs.reserve(args_spec_list.size() + 1); + (void)inputs.emplace_back(NewValueNode(prim::kPrimShard)); + for (size_t i = 0; i < args_spec_list.size(); ++i) { + (void)inputs.emplace_back(shard_fg->add_parameter()); + } + auto shard = shard_fg->NewCNodeInOrder(std::move(inputs)); + + FuncGraphPtr shard_child = nullptr; + { + TraceGuard guard(std::make_shared(shard_fg->debug_info())); + shard_child = GetShard(shard, origin_graph->parameters()); + } + shard_fg->set_output(NewValueNode(shard_child)); + return shard_fg; +} + +REGISTER_PYBIND_DEFINE(Shard_, ([](const py::module *m) { + (void)py::class_>(*m, "Shard_") + .def(py::init(), py::arg("fn")); + })); } // namespace prim } // namespace mindspore diff --git a/mindspore/ccsrc/frontend/operator/composite/composite.h b/mindspore/ccsrc/frontend/operator/composite/composite.h index 7af9f6e2d7f..2005a8eb6a1 100644 --- a/mindspore/ccsrc/frontend/operator/composite/composite.h +++ b/mindspore/ccsrc/frontend/operator/composite/composite.h @@ -213,6 +213,23 @@ class TupleGetItemTensor : public MetaFuncGraph { } }; using TupleGetItemTensorPtr = std::shared_ptr; + +class Shard : public MetaFuncGraph { + public: + explicit Shard(const string &name) : MetaFuncGraph(name) { + signatures_ = + // def shard(func:read, weight_list:read, in_axes:read, out_axes:read, device:read, level:read): + std::vector({{"func", SignatureEnumRW::kRWRead, SignatureEnumKind::kKindDefault}, + {"in_axes", SignatureEnumRW::kRWRead, SignatureEnumKind::kKindDefault}, + {"out_axes", SignatureEnumRW::kRWRead, SignatureEnumKind::kKindDefault}, + {"device", SignatureEnumRW::kRWRead, SignatureEnumKind::kKindDefault}, + {"level", SignatureEnumRW::kRWRead, SignatureEnumKind::kKindDefault}}); + } + ~Shard() override = default; + MS_DECLARE_PARENT(Shard, MetaFuncGraph) + + FuncGraphPtr GenerateFuncGraph(const AbstractBasePtrList &args_spec_list) override; +}; } // namespace prim } // namespace mindspore diff --git a/mindspore/ccsrc/frontend/operator/ops_front_infer_function.cc b/mindspore/ccsrc/frontend/operator/ops_front_infer_function.cc index 53894a200db..4420e662371 100644 --- a/mindspore/ccsrc/frontend/operator/ops_front_infer_function.cc +++ b/mindspore/ccsrc/frontend/operator/ops_front_infer_function.cc @@ -714,6 +714,26 @@ AbstractBasePtr InferImplJ(const AnalysisEnginePtr &, const PrimitivePtr &primit return AbstractFunction::MakeAbstractFunction(jv); } +AbstractBasePtr InferImplShard(const AnalysisEnginePtr &, const PrimitivePtr &primitive, + const AbstractBasePtrList &args_spec_list) { + // Inputs: func, in_axes, out_axes, device, level. + constexpr size_t shard_input_size = 5; + CheckArgsSize(primitive->name(), args_spec_list, shard_input_size); + MS_LOG(DEBUG) << "Evaluate Shard: " << args_spec_list[0]->ToString(); + + AbstractFunctionPtr x = dyn_cast(args_spec_list[0]); + MS_EXCEPTION_IF_NULL(x); + + AbstractFuncAtomPtrList shard_v; + auto build_shard_v = [&shard_v](const AbstractFuncAtomPtr &func) { + auto shard_closure = std::make_shared(func); + shard_v.push_back(shard_closure); + }; + x->Visit(build_shard_v); + + return AbstractFunction::MakeAbstractFunction(shard_v); +} + AbstractBasePtr InferImplFakeBprop(const AnalysisEnginePtr &, const PrimitivePtr &primitive, const AbstractBasePtrList &args_spec_list) { // Inputs: a tensor. @@ -779,6 +799,7 @@ REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(StringConcat, prim::kPrimStringConcat, InferI REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(DictLen, prim::kPrimDictLen, InferImplDictLen, nullptr); REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(FakeBprop, prim::kPrimFakeBprop, InferImplFakeBprop, nullptr); REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(J, prim::kPrimJ, InferImplJ, nullptr); +REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(Shard, prim::kPrimShard, InferImplShard, nullptr); REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(BroadcastGradientArgs, prim::kPrimBroadcastGradientArgs, InferImplBroadcastGradientArgs, nullptr); REGISTER_PRIMITIVE_FRONT_EVAL_IMPL(MakeSiice, prim::kPrimMakeSlice, InferImplMakeSlice, nullptr); diff --git a/mindspore/ccsrc/frontend/operator/ops_front_infer_function.h b/mindspore/ccsrc/frontend/operator/ops_front_infer_function.h index 2463413bb02..2d17812e3c3 100644 --- a/mindspore/ccsrc/frontend/operator/ops_front_infer_function.h +++ b/mindspore/ccsrc/frontend/operator/ops_front_infer_function.h @@ -55,6 +55,8 @@ AbstractBasePtr InferImplDictLen(const AnalysisEnginePtr &, const PrimitivePtr & const AbstractBasePtrList &args_spec_list); AbstractBasePtr InferImplJ(const AnalysisEnginePtr &, const PrimitivePtr &primitive, const AbstractBasePtrList &args_spec_list); +AbstractBasePtr InferImplShard(const AnalysisEnginePtr &, const PrimitivePtr &primitive, + const AbstractBasePtrList &args_spec_list); AbstractBasePtr InferImplFakeBprop(const AnalysisEnginePtr &, const PrimitivePtr &primitive, const AbstractBasePtrList &args_spec_list); AbstractBasePtr InferImplMakeRecord(const AnalysisEnginePtr &, const PrimitivePtr &primitive, diff --git a/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.cc b/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.cc index 1da4dbd7f91..d1e23866b8b 100644 --- a/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.cc +++ b/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.cc @@ -523,6 +523,28 @@ EvalResultPtr JEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &arg return res; } +EvalResultPtr ShardEvaluator::Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, + const AnfNodeConfigPtr &) { + AbstractBasePtrList args_spec_list; + (void)std::transform(args_conf_list.begin(), args_conf_list.end(), std::back_inserter(args_spec_list), + [](const ConfigPtr &conf) -> AbstractBasePtr { + MS_EXCEPTION_IF_NULL(conf); + return conf->ObtainEvalResult()->abstract(); + }); + MS_EXCEPTION_IF_NULL(evaluator_cache_mgr_); + auto eval_result = evaluator_cache_mgr_->GetValue(args_spec_list); + if (eval_result != nullptr) { + return eval_result; + } + + // Call the original evaluator, get the result: y = f(x) + EvalResultPtr result = evaluator_->Run(engine, args_conf_list, nullptr); + MS_EXCEPTION_IF_NULL(result); + auto res = std::make_shared(result->abstract(), std::make_shared()); + evaluator_cache_mgr_->SetValue(args_spec_list, res); + return res; +} + EvalResultPtr VirtualEvaluator::Eval(AnalysisEnginePtr, const AbstractBasePtrList &args_spec_list, const AnfNodeConfigPtr &out_conf) { if (args_spec_list.size() != args_spec_list_.size()) { diff --git a/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.h b/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.h index a5585000258..73d17b862e3 100644 --- a/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.h +++ b/mindspore/ccsrc/pipeline/jit/static_analysis/evaluator.h @@ -357,6 +357,41 @@ class JEvaluator : public Evaluator { AbstractFunctionPtr orig_func_; }; +class ShardEvaluator : public Evaluator { + public: + ShardEvaluator(const EvaluatorPtr &evaluator, const AbstractFunctionPtr &orig_func) + : Evaluator("ShardEvaluator"), evaluator_(evaluator), orig_func_(orig_func) {} + ~ShardEvaluator() override = default; + MS_DECLARE_PARENT(ShardEvaluator, Evaluator); + + AnfNodePtr bound_node() const override { + if (evaluator_ != nullptr) { + return evaluator_->bound_node(); + } + return bound_node_.lock(); + } + + void set_bound_node(const AnfNodePtr &node) override { + if (evaluator_ != nullptr) { + evaluator_->set_bound_node(node); + } + bound_node_ = AnfNodeWeakPtr(node); + } + + EvalResultPtr Eval(AnalysisEnginePtr, const AbstractBasePtrList &, const AnfNodeConfigPtr &) override { + MS_LOG(EXCEPTION) << "Should not be called, Run() method should be called"; + } + + EvalResultPtr Run(AnalysisEnginePtr engine, const ConfigPtrList &args_conf_list, + const AnfNodeConfigPtr &out_conf) override; + + std::string ToString() const override { return identifier_ + "_" + evaluator_->ToString(); } + + private: + EvaluatorPtr evaluator_; + AbstractFunctionPtr orig_func_; +}; + void BroadenArgs(const AbstractBasePtrList &args_spec_list, AbstractBasePtrList *broaded_args); } // namespace abstract } // namespace mindspore diff --git a/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.cc b/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.cc index 01b8a204502..4ac78e2ff1d 100644 --- a/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.cc +++ b/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.cc @@ -458,6 +458,14 @@ EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr &func) { + MS_EXCEPTION_IF_NULL(func); + AbstractFunctionPtr func_orig = func->fn(); + EvaluatorPtr evaluator_orig = GetEvaluatorFor(func_orig); + auto shard_evaluator = std::make_shared(evaluator_orig, func_orig); + return shard_evaluator; +} + EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const std::shared_ptr &func) { MS_EXCEPTION_IF_NULL(func); std::shared_ptr virtual_evaluator = @@ -495,6 +503,8 @@ EvaluatorPtr AnalysisEngine::_GetEvaluatorFor(const AbstractFunctionPtr &func) { return _GetEvaluatorFor(func->cast>()); } else if (func->isa()) { return _GetEvaluatorFor(func->cast>()); + } else if (func->isa()) { + return _GetEvaluatorFor(func->cast>()); } else if (func->isa()) { return _GetEvaluatorFor(func->cast>()); } else if (func->isa()) { diff --git a/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.h b/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.h index a847c3c4917..23a1ae31b41 100644 --- a/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.h +++ b/mindspore/ccsrc/pipeline/jit/static_analysis/static_analysis.h @@ -289,6 +289,7 @@ class AnalysisEngine : public std::enable_shared_from_this { EvaluatorPtr _GetEvaluatorFor(const std::shared_ptr &fn); EvaluatorPtr _GetEvaluatorFor(const std::shared_ptr &); EvaluatorPtr _GetEvaluatorFor(const std::shared_ptr &fn); + EvaluatorPtr _GetEvaluatorFor(const std::shared_ptr &fn); FuncGraphManagerPtr func_graph_manager() { return func_graph_manager_; } const AnfNodeConfigMap &anfnode_config_map() const { return anfnode_config_map_; } diff --git a/mindspore/core/abstract/abstract_function.cc b/mindspore/core/abstract/abstract_function.cc index a7007ea5e12..9d32ecc64e1 100644 --- a/mindspore/core/abstract/abstract_function.cc +++ b/mindspore/core/abstract/abstract_function.cc @@ -248,6 +248,20 @@ std::size_t JTransformedAbstractClosure::hash() const { return hash_value; } +bool ShardTransformedAbstractClosure::operator==(const AbstractFunction &other) const { + if (!other.isa()) { + return false; + } + auto other_transformed = static_cast(&other); + return fn_ == other_transformed->fn_; +} + +std::size_t ShardTransformedAbstractClosure::hash() const { + MS_EXCEPTION_IF_NULL(fn_); + auto hash_value = hash_combine(tid(), fn_->hash()); + return hash_value; +} + bool VirtualAbstractClosure::operator==(const AbstractFunction &other) const { if (!other.isa()) { return false; diff --git a/mindspore/core/abstract/abstract_function.h b/mindspore/core/abstract/abstract_function.h index 7d3a00c7f03..e1f455a498a 100644 --- a/mindspore/core/abstract/abstract_function.h +++ b/mindspore/core/abstract/abstract_function.h @@ -322,6 +322,36 @@ class MS_CORE_API JTransformedAbstractClosure final : public AbstractFuncAtom { AbstractFuncAtomPtr fn_; }; +/// \brief ShardTransformedAbstractClosure defines interface for abstract of Function +/// transformed through the application of Shard. +class MS_CORE_API ShardTransformedAbstractClosure final : public AbstractFuncAtom { + public: + /// \brief Constructor of ShardTransformedAbstractClosure + /// + /// \param[in] fn The AbstractFuncAtom transformed through the application of Shard. + explicit ShardTransformedAbstractClosure(const AbstractFuncAtomPtr &fn) : fn_(fn) {} + + /// \brief Destructor of ShardTransformedAbstractClosure + ~ShardTransformedAbstractClosure() override = default; + MS_DECLARE_PARENT(ShardTransformedAbstractClosure, AbstractFuncAtom) + + /// \brief Get the AbstractFuncAtom ShardTransformedAbstractClosure corresponding to. + /// + /// \return The AbstractFuncAtom ShardTransformedAbstractClosure corresponding to. + AbstractFuncAtomPtr fn() { return fn_; } + + AbstractFunctionPtr Copy() const override { return std::make_shared(fn_); } + + bool operator==(const AbstractFunction &other) const override; + + std::size_t hash() const override; + + std::string ToString() const override { return "Shard(" + fn_->ToString() + ")"; } + + private: + AbstractFuncAtomPtr fn_; +}; + /// \brief VirtualAbstractClosure defines interface for function with an explicitly /// fixed type signature. class MS_CORE_API VirtualAbstractClosure final : public AbstractFuncAtom { diff --git a/mindspore/core/base/core_ops.h b/mindspore/core/base/core_ops.h index 0a51499d06d..0bd82925cbc 100644 --- a/mindspore/core/base/core_ops.h +++ b/mindspore/core/base/core_ops.h @@ -685,6 +685,7 @@ inline const PrimitivePtr kPrimPyInterpret = std::make_shared("PyInte // Other primitive not used by backend but used in core; inline const PrimitivePtr kPrimStateSetItem = std::make_shared("state_setitem"); inline const PrimitivePtr kPrimJ = std::make_shared("J", kSideEffectPropagate); +inline const PrimitivePtr kPrimShard = std::make_shared("Shard", kSideEffectPropagate); // Used to build graph which have keyword arguments inline const PrimitivePtr kPrimExtractKeywordArg = std::make_shared("extract_keyword_arg"); diff --git a/mindspore/core/utils/trace_info.h b/mindspore/core/utils/trace_info.h index 0185bf0422c..c480a409b98 100644 --- a/mindspore/core/utils/trace_info.h +++ b/mindspore/core/utils/trace_info.h @@ -402,6 +402,14 @@ class TraceMixedPrecision : public TraceInfo { ~TraceMixedPrecision() override = default; TraceInfoPtr clone() override { return std::make_shared(*this); } }; + +class TraceShard : public TraceInfo { + public: + explicit TraceShard(const DebugInfoPtr &info) : TraceInfo(info) {} + ~TraceShard() override = default; + std::string name() const override { return "shard_ops"; } + TraceInfoPtr clone() override { return std::make_shared(*this); } +}; } // namespace mindspore #endif // MINDSPORE_CORE_UTILS_TRACE_INFO_H_ diff --git a/mindspore/ops/composite/__init__.py b/mindspore/ops/composite/__init__.py index c606c9ead58..227b250e06e 100644 --- a/mindspore/ops/composite/__init__.py +++ b/mindspore/ops/composite/__init__.py @@ -21,7 +21,7 @@ Pre-defined combination of operators. from .base import GradOperation, _Grad, HyperMap, Map, MultitypeFuncGraph, add_flags, \ - core, env_get, tail, zip_operation + core, env_get, tail, zip_operation, Shard from .clip_ops import clip_by_value, clip_by_global_norm from .multitype_ops.add_impl import hyper_add from .multitype_ops.ones_like_impl import ones_like @@ -60,4 +60,5 @@ __all__ = [ 'repeat_elements', 'sequence_mask', 'matmul', - '_Grad'] + '_Grad', + 'Shard'] diff --git a/mindspore/ops/composite/base.py b/mindspore/ops/composite/base.py index cb08a8eb27b..01655209f58 100644 --- a/mindspore/ops/composite/base.py +++ b/mindspore/ops/composite/base.py @@ -20,7 +20,7 @@ from functools import partial from types import FunctionType from mindspore import context -from ..._c_expression import EnvInstance_, GradOperation_, HyperMap_, Map_, MultitypeFuncGraph_, Tail_, \ +from ..._c_expression import EnvInstance_, GradOperation_, HyperMap_, Map_, MultitypeFuncGraph_, Tail_, Shard_, \ TupleAdd_, TupleSlice_, UnpackCall_, ZipOperation_, ListAppend_, TupleGetItemTensor_, ListInsert_ from ...common import dtype as mstype from ...common.api import ms_function, _pynative_executor, _wrap_func @@ -735,6 +735,48 @@ class Map(Map_): return tuple(map(func, *args_list)) +class Shard(Shard_): + """Shard operation""" + def __init__(self): + """Initialize Shard.""" + Shard_.__init__(self, 'Shard') + self.shard_fn = None + self.fn = None + self.in_axes = None + self.out_axes = None + self.device = None + self.level = None + + def __call__(self, fn, in_axes, out_axes, device, level=0): + if not isinstance(in_axes, tuple): + raise TypeError(f"For 'Shard', the 'in_axes' should be a tuple, but got {type(in_axes).__name__}") + if not isinstance(out_axes, tuple): + raise TypeError(f"For 'Shard', the 'out_axes' should be a tuple, " + f"but got {type(out_axes).__name__}") + if not isinstance(device, str): + raise TypeError(f"For 'Shard', the 'device' should be a string, " + f"but got {type(device).__name__}") + if not isinstance(level, int): + raise TypeError(f"For 'Shard', the 'level' should be an integer, " + f"but got {type(level).__name__}") + if self.shard_fn is not None and self.fn == fn and self.in_axes == in_axes and self.out_axes == out_axes and \ + self.device == device and self.level == level: + return self.shard_fn + shard_ = Shard() + + @ms_function + def after_shard(*args): + return shard_(fn, in_axes, out_axes, device, level)(*args) + + self.shard_fn = after_shard + self.fn = fn + self.in_axes = in_axes + self.out_axes = out_axes + self.device = device + self.level = level + return self.shard_fn + + class _ListAppend(ListAppend_): """ A metafuncgraph class that append one element to list. diff --git a/mindspore/ops/functional.py b/mindspore/ops/functional.py index b8ab5393a05..4a44ff688cd 100644 --- a/mindspore/ops/functional.py +++ b/mindspore/ops/functional.py @@ -28,7 +28,7 @@ from .primitive import Primitive from . import operations as P from .operations import _grad_ops from .operations import _csr_ops -from .composite import _Grad +from .composite import _Grad, Shard from .._c_expression import security typeof = Primitive('typeof') @@ -338,6 +338,10 @@ def vjp(fn, inputs, v): return wrap_container(*inputs, v) return wrap_container(inputs, v) +shard_fn = Shard() +def shard(fn, in_axes, out_axes, device, level=0): + return shard_fn(fn, in_axes, out_axes, device, level) + @constexpr def _raise_type_error(): diff --git a/tests/ut/cpp/operator/composite_test.cc b/tests/ut/cpp/operator/composite_test.cc index f81d4422894..30188e5c7f8 100644 --- a/tests/ut/cpp/operator/composite_test.cc +++ b/tests/ut/cpp/operator/composite_test.cc @@ -322,4 +322,52 @@ TEST_F(TestComposite, test_ZipOperation) { size_t expect = 3; ASSERT_EQ(real, expect); } + +/// Feature: Shard operation. +/// Description: Test the func_graph generation of Shard op and the inference of the Shard caller. +/// Expectation: Generate and the infer successfully. +TEST_F(TestComposite, test_shard) { + // Make origin func_graph which includes a relu node. + FuncGraphPtr origin_func_graph = std::make_shared(); + std::vector inputs; + inputs.push_back(NewValueNode(prim::kPrimRelu)); + inputs.push_back(origin_func_graph->add_parameter()); + CNodePtr relu = origin_func_graph->NewCNode(inputs); + inputs.clear(); + inputs.push_back(NewValueNode(prim::kPrimReturn)); + inputs.push_back(relu); + CNodePtr origin_return = origin_func_graph->NewCNode(inputs); + origin_func_graph->set_return(origin_return); + + // Make the func_graph which includes a Shard meta_func_graph. + FuncGraphPtr shard_func_graph = std::make_shared(); + MetaFuncGraphPtr shard_op = std::make_shared("shard_op"); + inputs.clear(); + inputs.push_back(NewValueNode(shard_op)); + inputs.push_back(NewValueNode(origin_func_graph)); + for (size_t i = 0; i < 4; ++i) { + inputs.push_back(NewValueNode(MakeValue(0))); + } + CNodePtr shard = shard_func_graph->NewCNode(inputs); + inputs.clear(); + inputs.push_back(shard); + inputs.push_back(shard_func_graph->add_parameter()); + CNodePtr shard_user = shard_func_graph->NewCNode(inputs); + inputs.clear(); + inputs.push_back(NewValueNode(prim::kPrimReturn)); + inputs.push_back(shard_user); + CNodePtr shard_return = shard_func_graph->NewCNode(inputs); + shard_func_graph->set_return(shard_return); + + auto tensor = UTCompositeUtils::ArrayInt32Of({2, 3, 4}); + AbstractBasePtrList args_spec_list = {tensor}; + + auto ret = engine_->Run(shard_func_graph, args_spec_list).inferred->abstract(); + ASSERT_NE(ret, nullptr); + ASSERT_TRUE(ret->isa()); + auto build_shape = ret->BuildShape(); + EXPECT_TRUE(build_shape->isa()); + auto shape = build_shape->cast(); + ASSERT_EQ(shape->shape(), std::vector({2, 3, 4})); +} } // namespace mindspore