From c5b2996b3fb9936b0b2ae2f1e9b658f9aeca332b Mon Sep 17 00:00:00 2001 From: zhangzhaoju Date: Mon, 27 Dec 2021 15:51:28 +0800 Subject: [PATCH] Resuse ConvertData of GenerateArgumentsKey --- mindspore/ccsrc/pipeline/jit/init.cc | 4 +- mindspore/ccsrc/pipeline/jit/pipeline.cc | 71 +++++++--------- mindspore/ccsrc/pipeline/jit/pipeline.h | 9 +- mindspore/python/mindspore/common/api.py | 85 +++++-------------- tests/ut/python/optimizer/test_python_pass.py | 36 ++++++-- 5 files changed, 86 insertions(+), 119 deletions(-) diff --git a/mindspore/ccsrc/pipeline/jit/init.cc b/mindspore/ccsrc/pipeline/jit/init.cc index bf459fbd4ca..b6ad89b3bec 100644 --- a/mindspore/ccsrc/pipeline/jit/init.cc +++ b/mindspore/ccsrc/pipeline/jit/init.cc @@ -115,11 +115,11 @@ PYBIND11_MODULE(_c_expression, m) { "Set values of weights.") .def("get_optimize_graph_proto", &GraphExecutorPy::GetOptimizeGraphProto, py::arg("phase") = py::str(""), "Get the optimize graph proto string.") - .def("set_jit_config", &GraphExecutorPy::SetJitConfig, py::arg("jit_config") = py::dict(), "Set the jit config."); + .def("set_jit_config", &GraphExecutorPy::SetJitConfig, py::arg("jit_config") = py::dict(), "Set the jit config.") + .def("generate_arguments_key", &GraphExecutorPy::GenerateArgumentsKey, "Generate unique key of argument."); (void)py::class_>(m, "EnvInstance_").def(py::init()); - (void)m.def("generate_arguments_key", &mindspore::pipeline::GenerateArgumentsKey, "Generate unique key of argument."); (void)m.def("real_run_op", &mindspore::pynative::RealRunOp, "Run op pynatively."); (void)m.def("reset_op_id", &mindspore::pipeline::ResetOpId, "Reset Operator Id"); (void)m.def("init_hccl", &mindspore::pipeline::InitHccl, "Init Hccl"); diff --git a/mindspore/ccsrc/pipeline/jit/pipeline.cc b/mindspore/ccsrc/pipeline/jit/pipeline.cc index 2a1e9b4ec06..5ec59d7b43f 100644 --- a/mindspore/ccsrc/pipeline/jit/pipeline.cc +++ b/mindspore/ccsrc/pipeline/jit/pipeline.cc @@ -476,33 +476,39 @@ void CheckArgsValid(const py::tuple &args) { } } -py::object GenerateArgumentsKey(const std::unordered_map &args, bool enable_tuple_broaden) { +py::object GraphExecutorPy::GenerateArgumentsKey(const py::tuple &args, bool enable_tuple_broaden) { MS_LOG(DEBUG) << "GenerateArgumentsKey args size:" << args.size(); + abstract::AbstractBasePtrList args_spec; - - for (const auto &arg : args) { - if (py::isinstance(arg.second)) { - MS_LOG(EXCEPTION) << "GenerateArgumentsKey failed, argument input should not be py::module"; - } + cur_convert_input_.clear(); + std::size_t size = args.size(); + for (std::size_t i = 0; i < size; i++) { ValuePtr converted = nullptr; - if (!parse::ConvertData(arg.second, &converted)) { - MS_LOG(EXCEPTION) << "GenerateArgumentsKey convert arg failed"; + if (!parse::ConvertData(args[i], &converted)) { + MS_EXCEPTION(TypeError) + << "The inputs types of the outermost network support bool, int, float, None, tensor, " + "mstype.Number(mstype.bool, mstype.int, mstype.float, mstype.uint), " + "and tuple or list containing only these types, and dict whose values are these types, but the " + << i << "th arg type is " << args[i].get_type() << ", value is '" << py::str(args[i]) << "'."; } - args_spec.push_back(ArgsToAbstract(converted, enable_tuple_broaden)); + AbstractBasePtr ptr = ArgsToAbstract(converted, enable_tuple_broaden); + args_spec.push_back(ptr); + cur_convert_input_.emplace(args[i].ptr(), ptr); } - uint64_t key; + // If cache matched no need CheckArgsValid auto iter = g_args_cache.find(args_spec); - if (iter == g_args_cache.end()) { - static uint64_t key_counter = 0; - key = key_counter; - ++key_counter; - g_args_cache[args_spec] = key; - MS_LOG(INFO) << "Generate a new compile key for new args, key: " << key; - } else { - key = iter->second; + if (iter != g_args_cache.end()) { + return py::int_(iter->second); } - return py::int_(key); + + // Check if the args of function or net is valid. + CheckArgsValid(args); + + static uint64_t key_counter = 0; + g_args_cache[args_spec] = key_counter; + MS_LOG(INFO) << "Generate a new compile key for new args, key: " << key_counter; + return py::int_(key_counter++); } py::bool_ VerifyInputSignature(const py::list &input_signature, const py::tuple &inputs) { @@ -981,13 +987,6 @@ bool GraphExecutorPy::CompileInner(const py::object &source_obj, const py::tuple MS_LOG(ERROR) << "The `phase` must be string."; return false; } - // Check if the function or net is valid. - if (py::isinstance(source_obj)) { - MS_LOG(ERROR) << "The source object to compile should not be None."; - return false; - } - // Check if the args of function or net is valid. - CheckArgsValid(args); auto phase = py::cast(phase_obj); MS_LOG(INFO) << "Start compiling, phase: " << phase; @@ -1023,6 +1022,13 @@ bool GraphExecutorPy::CompileInner(const py::object &source_obj, const py::tuple std::size_t size = args.size(); for (std::size_t i = 0; i < size; i++) { ValuePtr converted = nullptr; + // In some parallel mode need full_tensor which cause the args of GenerateArgumentsKey not same to compile, + // So can't use cur_convert_input_ directly. + auto iter = cur_convert_input_.find(args[i].ptr()); + if (iter != cur_convert_input_.end()) { + args_spec.push_back(iter->second); + continue; + } bool succ = parse::ConvertData(args[i], &converted); if (!succ) { MS_LOG(EXCEPTION) << "Fail to convert the " << i << "th argument, args[" << i << "]: " << py::str(args[i]); @@ -1290,10 +1296,6 @@ void ProcessVmArgInner(const py::tuple &args, const ResourcePtr &res, VectorRef bool arg_list_inited = !arg_list->empty(); for (std::size_t i = 0; i < size; i++) { py::object arg = args[i]; - auto ms_context = MsContext::GetInstance(); - if (ms_context->backend_policy() == kMsConvert && py::isinstance(arg)) { - MS_LOG(EXCEPTION) << "The " << i << "th arg is numpy array, not tensor."; - } ValuePtr converted = nullptr; bool succ = parse::ConvertData(arg, &converted); if (!succ) { @@ -1369,15 +1371,6 @@ py::object GraphExecutorPy::Run(const py::tuple &args, const py::object &phase_o auto ret_val = std::make_shared(); if (info_.count(phase) != 0 && info_[phase]->func_graph != nullptr) { if (IsGraphOutputValueNodeOrParameter(info_[phase]->func_graph->output(), args, ret_val)) { - // Check the input arg must be Tensor when backend is "ms". - if (MsContext::GetInstance()->backend_policy() == kMsConvert) { - for (std::size_t i = 0; i < size; i++) { - ValuePtr converted = nullptr; - if (!parse::ConvertData(args[i], &converted)) { - MS_LOG(EXCEPTION) << "The " << i << "th arg convert failed."; - } - } - } return *ret_val; } } diff --git a/mindspore/ccsrc/pipeline/jit/pipeline.h b/mindspore/ccsrc/pipeline/jit/pipeline.h index 208701dac8d..eecb62d801a 100644 --- a/mindspore/ccsrc/pipeline/jit/pipeline.h +++ b/mindspore/ccsrc/pipeline/jit/pipeline.h @@ -24,6 +24,7 @@ #include #include #include +#include #include "pybind11/pybind11.h" @@ -128,9 +129,11 @@ class GraphExecutorPy : public std::enable_shared_from_this { std::map> FetchInfoForQuantExport( const std::string &phase); + // Generate a key for mapping function graph + py::object GenerateArgumentsKey(const py::tuple &args, bool enable_tuple_broaden = false); + private: GraphExecutorPy(); - void ConvertObjectToTensors(const py::dict &dict, std::map *tensors); void GetWeightInfo(const CNodePtr &root_node, const AnfNodePtr &weight_node, std::map> *fake_quant_table); void GetGeBackendPolicy() const; @@ -153,15 +156,13 @@ class GraphExecutorPy : public std::enable_shared_from_this { bool enable_tuple_broaden_{false}; py::list compile_cache_dep_files_; py::dict weights_; + std::map cur_convert_input_; }; using GraphExecutorPyPtr = std::shared_ptr; std::string GetJitLevel(); void CheckArgsValid(const py::tuple &args); -// Generate a key for mapping function graph -py::object GenerateArgumentsKey(const std::unordered_map &args, - bool enable_tuple_broaden = false); py::bool_ VerifyInputSignature(const py::list &input_signature, const py::tuple &inputs); bool InitDistribute(const std::map &options); diff --git a/mindspore/python/mindspore/common/api.py b/mindspore/python/mindspore/common/api.py index 8d09425b791..14bc358672e 100644 --- a/mindspore/python/mindspore/common/api.py +++ b/mindspore/python/mindspore/common/api.py @@ -29,7 +29,7 @@ from mindspore import log as logger from mindspore._extends.remote import kernel_build_server from .tensor import Tensor as MsTensor from .tensor import CSRTensor as MsCSRTensor -from .._c_expression import generate_arguments_key, GraphExecutor_, Tensor, MetaTensor, CSRTensor, PynativeExecutor_ +from .._c_expression import GraphExecutor_, Tensor, MetaTensor, CSRTensor, PynativeExecutor_ from .._c_expression import verify_inputs_signature, init_exec_dataset, _set_dataset_mode_config, init_pipeline from ..parallel._ps_context import _is_role_pserver, _is_role_sched from ..parallel._utils import _get_device_num, _get_global_rank, _need_to_full, _check_full_batch, _to_full_tensor, \ @@ -42,30 +42,6 @@ ms_compile_cache = {} BROADCAST_PHASE = "_broadcast_" -def _convert_function_arguments(fn, *args): - """ - Process the fn default parameters. - - Args: - fn (Function): The function to be parsed. - args (tuple): The parameters of the function. - """ - arguments_dict = OrderedDict() - parse_method = None - if isinstance(fn, (types.FunctionType, types.MethodType)): - parse_method = fn.__name__ - index = 0 - for value in args: - arguments_dict[f'arg{index}'] = value - index = index + 1 - logger.debug("fn(%r) full parameters dict is: %r", fn, arguments_dict) - converted = True - else: - logger.warning("Find error: fn isn't function or method") - converted = False - return converted, arguments_dict, parse_method - - def _wrap_func(fn): """ Wrapper function, convert return data to tensor or tuple of tensor. @@ -125,6 +101,8 @@ if sys.argv and sys.argv[0] != '': entry_script_path_dir = os.path.split(entry_script_path)[0] if entry_script_path_dir in sys_path: sys_path.remove(entry_script_path_dir) + + def _in_sys_path(file_path): for path in sys_path: if file_path.startswith(path): @@ -207,10 +185,14 @@ class _MindsporeFunctionExecutor: """ def __init__(self, fn, ms_create_time, input_signature=None, obj=None): + init_pipeline() + if not isinstance(fn, (types.FunctionType, types.MethodType)): + raise RuntimeError('fn {} is not function or method'.format(fn)) + self.fn = fn self.input_signature = input_signature self.obj = None - if hasattr(obj, fn.__name__): + if obj and hasattr(obj, fn.__name__): self.obj = obj self._graph_executor = GraphExecutor_.get_instance() self._create_time = ms_create_time @@ -227,7 +209,7 @@ class _MindsporeFunctionExecutor: init_phase = "init_subgraph" + graph_name[graph_name.find("."):] _exec_init_graph(self.obj, init_phase) - def compile(self, args_list, arg_names, method_name): + def compile(self, args_list, method_name): """Returns pipeline for the given args.""" # Verify the signature for both function and method if self.input_signature is not None: @@ -240,9 +222,8 @@ class _MindsporeFunctionExecutor: if not is_valid_input: raise ValueError("Inputs is incompatible with input signature!") - dic = dict(zip(arg_names, args_list)) generate_name = self.fn.__module__ + "." + self.fn.__name__ + "." + self.fn.__code__.co_filename + "." + \ - str(self.fn.__code__.co_firstlineno) + '.' + str(id(self.fn)) + str(self.fn.__code__.co_firstlineno) + '.' + str(id(self.fn)) if _pynative_executor.grad_flag(): generate_name = generate_name + ".grad" self.fn.__parse_method__ = method_name @@ -263,7 +244,7 @@ class _MindsporeFunctionExecutor: else: self.enable_tuple_broaden = False self._graph_executor.set_enable_tuple_broaden(self.enable_tuple_broaden) - key = generate_arguments_key(dic, self.enable_tuple_broaden) + key = self._graph_executor.generate_arguments_key(args_list, self.enable_tuple_broaden) phase = generate_name + '.' + str(key) if phase in ms_compile_cache.keys(): return phase @@ -282,18 +263,11 @@ class _MindsporeFunctionExecutor: @_wrap_func def __call__(self, *args): - init_pipeline() - converted, arguments_dict, parse_method = _convert_function_arguments(self.fn, *args) - if not converted: - raise RuntimeError('Process function parameter is failure') - - args_list = tuple(arguments_dict.values()) - arg_names = tuple(arguments_dict.keys()) + args_list = args if self.obj is not None: args_list = args_list[1:] - arg_names = arg_names[1:] - phase = self.compile(args_list, arg_names, parse_method) + phase = self.compile(args_list, self.fn.__name__) if context.get_context("precompile_only"): return None @@ -390,21 +364,6 @@ def ms_function(fn=None, obj=None, input_signature=None): return wrap_mindspore -def _generate_pip_args(obj, *args, method="construct"): - """Generate arguments for pipeline.""" - if hasattr(obj, method): - fn = getattr(obj, method) - else: - raise AttributeError('The process method is not exist') - converted, arguments_dict, parse_method = _convert_function_arguments(fn, *args) - if not converted: - raise RuntimeError('Process method parameter is failure') - args_list = tuple(arguments_dict.values()) - args_names = tuple(arguments_dict.keys()) - obj.__parse_method__ = parse_method - return args_names, args_list - - def _get_auto_split_param_names(parameter_layout_dict): auto_split_param_names = [] for key, value in parameter_layout_dict.items(): @@ -630,15 +589,14 @@ class _CellGraphExecutor: Str, the full phase of the cell. Bool, if the graph has been compiled before, return False, else return True. """ - - args_names, args_list = _generate_pip_args(obj, *args) - dic = dict(zip(args_names, args_list)) + obj.__parse_method__ = obj.construct.__name__ + args_list = args if hasattr(obj, "enable_tuple_broaden"): self.enable_tuple_broaden = obj.enable_tuple_broaden else: self.enable_tuple_broaden = False self._graph_executor.set_enable_tuple_broaden(self.enable_tuple_broaden) - key = generate_arguments_key(dic, self.enable_tuple_broaden) + key = self._graph_executor.generate_arguments_key(args_list, self.enable_tuple_broaden) obj.arguments_key = str(key) phase = phase + '.' + str(obj.create_time) + '.' + str(id(obj)) + '.' + obj.arguments_key @@ -653,8 +611,7 @@ class _CellGraphExecutor: is_sink_mode = args and isinstance(args[0], Tensor) and args[0].virtual_flag if auto_parallel_mode and _need_to_full() and not is_sink_mode and obj.auto_parallel_compile_and_run(): - args_full = _to_full_tensor(args, _get_device_num(), _get_global_rank()) - _, args_list = _generate_pip_args(obj, *args_full) + args_list = _to_full_tensor(args, _get_device_num(), _get_global_rank()) enable_ge = context.get_context("enable_ge") self._graph_executor.set_weights_values(obj.parameters_dict()) @@ -743,12 +700,8 @@ class _CellGraphExecutor: def _exec_pip(self, obj, *args, phase=''): """Execute the generated pipeline.""" fn = obj.construct - converted, arguments_dict, parse_method = _convert_function_arguments(fn, *args) - if not converted: - raise RuntimeError('Process method parameter is failure') - args_list = tuple(arguments_dict.values()) - obj.__parse_method__ = parse_method - return self._graph_executor(args_list, phase) + obj.__parse_method__ = fn.__name__ + return self._graph_executor(args, phase) def run(self, obj, *args, phase='predict'): """ diff --git a/tests/ut/python/optimizer/test_python_pass.py b/tests/ut/python/optimizer/test_python_pass.py index 811f71777d9..ba0dc1c5e92 100644 --- a/tests/ut/python/optimizer/test_python_pass.py +++ b/tests/ut/python/optimizer/test_python_pass.py @@ -20,24 +20,23 @@ from mindspore import context from mindspore.common.tensor import Tensor from mindspore.ops import operations as P from mindspore.ops import _constants as Constants -from mindspore.graph_utils.python_pass import register_pass, unregister_pass, set_renorm, gen_new_parameter,\ +from mindspore.graph_utils.python_pass import register_pass, unregister_pass, set_renorm, gen_new_parameter, \ cancel_new_parameter, set_reopt -from mindspore.common.api import _generate_pip_args -from mindspore._c_expression import generate_arguments_key, GraphExecutor_ +from mindspore._c_expression import GraphExecutor_ from mindspore.graph_utils.graph_pattern import OneOf, Prim, Call, NoneOf, Any, NewTensor, NewParameter, Imm context.set_context(mode=context.GRAPH_MODE) + def get_func_graph(obj, *args, phase="validate"): - args_names, args_list = _generate_pip_args(obj, *args) - dic = dict(zip(args_names, args_list)) - key = generate_arguments_key(dic, False) + _executor = GraphExecutor_.get_instance() + key = _executor.generate_arguments_key(args, False) obj.arguments_key = str(key) phase = phase + '.' + str(obj.create_time) + '.' + str(id(obj)) + '.' + obj.arguments_key - _executor = GraphExecutor_.get_instance() - _executor.compile(obj, args_list, phase, False) + _executor.compile(obj, args, phase, False) return _executor.get_func_graph(phase) + def test_softmax_relu(): """ Use python pass to transform from Softmax to ReLU. @@ -57,6 +56,7 @@ def test_softmax_relu(): assert "ReLU" in transformed_repr assert "Softmax" not in transformed_repr + def test_prim(): inputs = Tensor(np.ones([42]), mindspore.float16) softmax_model = nn.Softmax() @@ -74,6 +74,7 @@ def test_prim(): assert "ReLU" in transformed_repr assert "Softmax" not in transformed_repr + def test_softmax_relu_sigmoid(): """ Use python pass to transform from Softmax(x) to ReLU(Sigmoid(x)). @@ -121,11 +122,13 @@ def test_isin_pattern_0(): relu6_pattern = Prim(P.ReLU6()) target = Call(relu6_pattern, [x]) return pattern, target + transformed_repr = get_func_graph(softmax_model, inputs).get_return().expanded_str(2) unregister_pass(softmax_relu_pass) assert "ReLU6" in transformed_repr assert "Softmax" not in transformed_repr + def test_isin_pattern_1(): """ Test IsIn. IsIn is used as nested inputs for the target in this case. @@ -145,11 +148,13 @@ def test_isin_pattern_1(): neg_ops = Prim(P.Neg()) target = Call(neg_ops, [pattern]) return pattern, target + transformed_repr = get_func_graph(softmax_model, inputs).get_return().expanded_str(4) unregister_pass(softmax_neg_pass) assert "Neg" in transformed_repr assert "Softmax" in transformed_repr + def test_isnot_pattern_0(): """ Test IsNot pattern which expresses the IsNot semantics. @@ -157,6 +162,7 @@ def test_isnot_pattern_0(): """ set_renorm(False) set_reopt(False) + class ConvBN(nn.Cell): def __init__(self): super(ConvBN, self).__init__() @@ -167,10 +173,12 @@ def test_isnot_pattern_0(): self.mean = Tensor(np.ones([32]), mindspore.float32) self.variance = Tensor(np.ones([32]), mindspore.float32) self.bn = P.BatchNorm() + def construct(self, x): x = self.conv(x, self.conv_weight) x = self.bn(x, self.scale, self.bias, self.mean, self.variance) return x + inputs = Tensor(np.random.normal(0, 1, (10, 32, 32, 32)), mindspore.float32) conv_bn_model = ConvBN() @@ -202,6 +210,7 @@ def test_isnot_pattern_0(): assert "Softmax" in transformed_repr set_renorm(True) + def test_isnot_pattern_1(): """ Test IsNot pattern which expresses the IsNot semantics. @@ -228,6 +237,7 @@ def test_isnot_pattern_1(): assert "ReLU6" in transformed_repr assert "Softmax" not in transformed_repr + def test_newtensor_pattern(): """ Test NewTensor pattern in the target @@ -246,12 +256,14 @@ def test_newtensor_pattern(): new_weight = NewTensor(weight_tensor) target = Call(P.AddN(), [x, new_weight]) return pattern, target + transformed_repr = get_func_graph(softmax_model, inputs).get_return().expanded_str(2) unregister_pass(softmax_addn_pass) assert "AddN" in transformed_repr assert "Softmax" not in transformed_repr set_renorm(True) + def test_newparameter_pattern(): """ Test NewParameter pattern in the target @@ -261,6 +273,7 @@ def test_newparameter_pattern(): set_renorm(False) set_reopt(False) + @register_pass(requires_grad=False, run_only_once=True) def softmax_addn_pass(): x = Any() @@ -273,12 +286,14 @@ def test_newparameter_pattern(): target_0 = Call(P.MatMul(), [new_para_0, new_para_1]) target = Call("MakeTuple", [target_0]) return pattern, target + transformed_repr = get_func_graph(softmax_model, inputs).get_return().expanded_str(5) unregister_pass(softmax_addn_pass) assert "MatMul" in transformed_repr assert "MakeTuple" in transformed_repr assert "Softmax" not in transformed_repr + def test_imm_target(): """ Test NewParameter pattern in the target @@ -288,6 +303,7 @@ def test_imm_target(): set_renorm(False) set_reopt(False) + @register_pass(requires_grad=False, run_only_once=True) def softmax_pass(): x = Any() @@ -296,12 +312,14 @@ def test_imm_target(): target_0 = Call("MakeTuple", [pattern]) target = Call(Constants.kTupleGetItem, [target_0, imm]) return pattern, target + transformed_repr = get_func_graph(softmax_model, inputs).get_return().expanded_str(5) unregister_pass(softmax_pass) assert "MakeTuple" in transformed_repr assert Constants.kTupleGetItem in transformed_repr assert "Softmax" in transformed_repr + def test_gen_new_parameter(): """ Test gen_new_parameter @@ -314,6 +332,7 @@ def test_gen_new_parameter(): set_renorm(False) set_reopt(False) gen_new_parameter(new_para) + @register_pass(requires_grad=False, run_only_once=True) def softmax_make_tuple_pass(): x = Any() @@ -322,6 +341,7 @@ def test_gen_new_parameter(): target = Call("MakeTuple", [pattern, new_para]) return pattern, target + transformed_repr = get_func_graph(softmax_model, inputs).get_return().expanded_str(5) assert "Merlin" in transformed_repr unregister_pass(softmax_make_tuple_pass)