diff --git a/model.py b/model.py index b71e59171f76082f4f9457080e965557c6695be2..28a6befaa8027c9272eda126cab89ac8134d6903 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,6 +55,7 @@ class _StepSync(Callback): def step_end(run_context): _pynative_executor.sync() +# what we can modify class Model: """ High-Level API for training or inference. @@ -132,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 @@ -158,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): @@ -171,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 @@ -180,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( @@ -198,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 = "" @@ -1081,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"] diff --git a/test.xml b/test.xml new file mode 100644 index 0000000000000000000000000000000000000000..e8aecca4860997a20b09780e1e97b4606295a79c --- /dev/null +++ b/test.xml @@ -0,0 +1 @@ + 123 \ No newline at end of file