From a727fe2ac89cf341228e2268902a92a1d2f2a57e Mon Sep 17 00:00:00 2001 From: liutongtong Date: Fri, 19 Nov 2021 11:33:43 +0800 Subject: [PATCH] codecheck warning clean --- mindspore/common/api.py | 1 + mindspore/nn/cell.py | 34 ++++++++++++++++++++-------------- 2 files changed, 21 insertions(+), 14 deletions(-) diff --git a/mindspore/common/api.py b/mindspore/common/api.py index 50192dd38a1..5716d7bdd46 100644 --- a/mindspore/common/api.py +++ b/mindspore/common/api.py @@ -281,6 +281,7 @@ def ms_function(fn=None, obj=None, input_signature=None): def wrap_mindspore(func): ms_create_time = int(time.time() * 1e9) + @wraps(func) def staging_specialize(*args): if obj is not None: diff --git a/mindspore/nn/cell.py b/mindspore/nn/cell.py index 8eca8dec80f..94d6d042586 100755 --- a/mindspore/nn/cell.py +++ b/mindspore/nn/cell.py @@ -548,6 +548,24 @@ class Cell(Cell_): if hasattr(self, '_cell_init_args'): self.cell_init_args += str({name: value}) + def _set_attr_for_params(self, name, value): + if isinstance(value, Tensor) and self._params[name] is not None: + self._params[name].set_data(value) + elif value is not None: + raise TypeError(f"The type of value should be Parameter or ParameterTuple, " + f"but got {type(value).__name__}.") + else: + self.insert_param_to_cell(name, None) + + def _set_attr_for_tensor(self, name, value): + if context.get_context("mode") == context.PYNATIVE_MODE: + tensor_list = self.__dict__.get('_tensor_list') + if name in self.__dict__: + del self.__dict__[name] + tensor_list[name] = value + else: + object.__setattr__(self, name, value) + def _check_param_list_tuple(self, value): """ Check the type of input in list or tuple is Parameter. @@ -562,7 +580,6 @@ class Cell(Cell_): def __setattr__(self, name, value): cells = self.__dict__.get('_cells') params = self.__dict__.get('_params') - tensor_list = self.__dict__.get('_tensor_list') if isinstance(value, Parameter): self._set_attr_for_parameter(name, value) elif isinstance(value, ParameterTuple): @@ -572,24 +589,13 @@ class Cell(Cell_): elif isinstance(value, Cell): self._set_attr_for_cell(name, value) elif params and name in params: - if isinstance(value, Tensor) and self._params[name] is not None: - self._params[name].set_data(value) - elif value is not None: - raise TypeError(f"The type of value should be Parameter or ParameterTuple, " - f"but got {type(value).__name__}.") - else: - self.insert_param_to_cell(name, None) + self._set_attr_for_params(name, value) elif cells and name in cells: if value is not None: raise TypeError(f"The type of value should be cell, but got {type(value).__name__}.") self._cells[name] = None elif isinstance(value, Tensor): - if context.get_context("mode") == context.PYNATIVE_MODE: - if name in self.__dict__: - del self.__dict__[name] - tensor_list[name] = value - else: - object.__setattr__(self, name, value) + self._set_attr_for_tensor(name, value) else: if isinstance(value, Primitive): value.set_prim_instance_name(name)