From fd0db0dec11e6e0a9e02c5f8352d95a07bc1975a Mon Sep 17 00:00:00 2001 From: zhangyx95 <9473099+zhangyx95@user.noreply.gitee.com> Date: Tue, 15 Feb 2022 03:09:05 +0000 Subject: [PATCH 1/3] update model.py. --- model.py | 1 + 1 file changed, 1 insertion(+) diff --git a/model.py b/model.py index b71e591..ed47cb5 100644 --- a/model.py +++ b/model.py @@ -55,6 +55,7 @@ class _StepSync(Callback): def step_end(run_context): _pynative_executor.sync() +# we can do it class Model: """ High-Level API for training or inference. -- Gitee From f6e1149e9c3f4905abf444825760fb8b3f24c3d1 Mon Sep 17 00:00:00 2001 From: zhangyx95 <9473099+zhangyx95@user.noreply.gitee.com> Date: Tue, 15 Feb 2022 06:27:18 +0000 Subject: [PATCH 2/3] update model.py. --- model.py | 19 +++++-------------- 1 file changed, 5 insertions(+), 14 deletions(-) diff --git a/model.py b/model.py index ed47cb5..28a6bef 100644 --- a/model.py +++ b/model.py @@ -39,7 +39,7 @@ from .dataset_helper import DatasetHelper, connect_network_with_dataset from . import amp from ..common.api import _pynative_executor - +# test def _transfer_tensor_to_tuple(inputs): """ If the input is a tensor, convert it to a tuple. If not, the output is unchanged. @@ -55,7 +55,7 @@ class _StepSync(Callback): def step_end(run_context): _pynative_executor.sync() -# we can do it +# what we can modify class Model: """ High-Level API for training or inference. @@ -133,7 +133,7 @@ class Model: >>> dataset = create_custom_dataset() >>> model.train(2, dataset) """ - +# what we can modify def __init__(self, network, loss_fn=None, optimizer=None, metrics=None, eval_network=None, eval_indexes=None, amp_level="O0", boost_level="O0", **kwargs): self._network = network @@ -159,7 +159,7 @@ class Model: self._train_network = self._build_train_network() self._build_eval_network(metrics, self._eval_network, eval_indexes) self._build_predict_network() - +# what we can modify def _check_for_graph_cell(self, kwargs): """Check for graph cell""" if not isinstance(self._network, nn.GraphCell): @@ -172,7 +172,6 @@ class Model: "but got 'loss_fn': {}, 'optimizer': {}.".format(self._loss_fn, self._optimizer)) if kwargs: raise ValueError("For 'Model', the '**kwargs' argument should be empty when network is a GraphCell.") - def _process_amp_args(self, kwargs): if self._amp_level in ["O0", "O3"]: self._keep_bn_fp32 = False @@ -181,7 +180,7 @@ class Model: if 'loss_scale_manager' in kwargs: self._loss_scale_manager = kwargs['loss_scale_manager'] self._loss_scale_manager_set = True - +# what we can modify def _check_amp_level_arg(self, optimizer, amp_level): if optimizer is None and amp_level != "O0": raise ValueError( @@ -199,7 +198,6 @@ class Model: dataset.__model_hash__ = hash(self) if hasattr(dataset, '__model_hash__') and dataset.__model_hash__ != hash(self): raise RuntimeError('The Dataset cannot be bound to different models, please create a new dataset.') - def _build_boost_network(self, kwargs): """Build the boost network.""" boost_config_dict = "" @@ -1082,35 +1080,28 @@ class Model: raise RuntimeError('Infer predict layout only supports semi auto parallel and auto parallel mode.') _parallel_predict_check() check_input_data(*predict_data, data_class=Tensor) - predict_net = self._predict_network # Unlike the cases in build_train_network() and build_eval_network(), 'multi_subgraphs' is not set predict_net.set_auto_parallel() predict_net.set_train(False) predict_net.compile(*predict_data) return predict_net.parameter_layout_dict - def _flush_from_cache(self, cb_params): """Flush cache data to host if tensor is cache enable.""" params = cb_params.train_network.get_parameters() for param in params: if param.cache_enable: Tensor(param).flush_from_cache() - @property def train_network(self): """Get the model's train_network.""" return self._train_network - @property def predict_network(self): """Get the model's predict_network.""" return self._predict_network - @property def eval_network(self): """Get the model's eval_network.""" return self._eval_network - - __all__ = ["Model"] -- Gitee From b19cf211470cb6841cd5f3340e62db74b61849b2 Mon Sep 17 00:00:00 2001 From: zhangyx95 <9473099+zhangyx95@user.noreply.gitee.com> Date: Fri, 22 Apr 2022 03:15:56 +0000 Subject: [PATCH 3/3] add test.xml. --- test.xml | 1 + 1 file changed, 1 insertion(+) create mode 100644 test.xml diff --git a/test.xml b/test.xml new file mode 100644 index 0000000..e8aecca --- /dev/null +++ b/test.xml @@ -0,0 +1 @@ + 123 \ No newline at end of file -- Gitee