diff --git a/deepmd/fit/dipole.py b/deepmd/fit/dipole.py index 383ea17f1f..2935aa2b06 100644 --- a/deepmd/fit/dipole.py +++ b/deepmd/fit/dipole.py @@ -182,7 +182,7 @@ def init_variables(self, suffix : str suffix to name scope """ - self.fitting_net_variables = get_fitting_net_variables_from_graph_def(graph_def) + self.fitting_net_variables = get_fitting_net_variables_from_graph_def(graph_def, suffix=suffix) def enable_mixed_precision(self, mixed_prec : dict = None) -> None: @@ -195,4 +195,4 @@ def enable_mixed_precision(self, mixed_prec : dict = None) -> None: The mixed precision setting used in the embedding net """ self.mixed_prec = mixed_prec - self.fitting_precision = get_precision(mixed_prec['output_prec']) \ No newline at end of file + self.fitting_precision = get_precision(mixed_prec['output_prec']) diff --git a/deepmd/fit/ener.py b/deepmd/fit/ener.py index 61d70045d8..91ec662ce4 100644 --- a/deepmd/fit/ener.py +++ b/deepmd/fit/ener.py @@ -538,7 +538,7 @@ def init_variables(self, suffix : str suffix to name scope """ - self.fitting_net_variables = get_fitting_net_variables_from_graph_def(graph_def) + self.fitting_net_variables = get_fitting_net_variables_from_graph_def(graph_def, suffix=suffix) if self.numb_fparam > 0: self.fparam_avg = get_tensor_by_name_from_graph(graph, 'fitting_attr%s/t_fparam_avg' % suffix) self.fparam_inv_std = get_tensor_by_name_from_graph(graph, 'fitting_attr%s/t_fparam_istd' % suffix) diff --git a/deepmd/fit/polar.py b/deepmd/fit/polar.py index 3f1b7daa6b..3bb3d9966b 100644 --- a/deepmd/fit/polar.py +++ b/deepmd/fit/polar.py @@ -389,7 +389,7 @@ def init_variables(self, suffix : str suffix to name scope """ - self.fitting_net_variables = get_fitting_net_variables_from_graph_def(graph_def) + self.fitting_net_variables = get_fitting_net_variables_from_graph_def(graph_def, suffix=suffix) def enable_mixed_precision(self, mixed_prec : dict = None) -> None: diff --git a/deepmd/utils/graph.py b/deepmd/utils/graph.py index fafca75f20..10e13730c5 100644 --- a/deepmd/utils/graph.py +++ b/deepmd/utils/graph.py @@ -238,7 +238,7 @@ def get_embedding_net_variables(model_file : str, suffix: str = "") -> Dict: return get_embedding_net_variables_from_graph_def(graph_def, suffix=suffix) -def get_fitting_net_nodes_from_graph_def(graph_def: tf.GraphDef) -> Dict: +def get_fitting_net_nodes_from_graph_def(graph_def: tf.GraphDef, suffix: str = "") -> Dict: """ Get the fitting net nodes with the given tf.GraphDef object @@ -246,13 +246,22 @@ def get_fitting_net_nodes_from_graph_def(graph_def: tf.GraphDef) -> Dict: ---------- graph_def The input tf.GraphDef object + suffix + suffix of the scope Returns ---------- Dict The fitting net nodes within the given tf.GraphDef object """ - fitting_net_nodes = get_pattern_nodes_from_graph_def(graph_def, FITTING_NET_PATTERN) + if suffix != "": + fitting_net_pattern = FITTING_NET_PATTERN\ + .replace('/idt', suffix + '/idt')\ + .replace('/bias', suffix + '/bias')\ + .replace('/matrix', suffix + '/matrix') + else: + fitting_net_pattern = FITTING_NET_PATTERN + fitting_net_nodes = get_pattern_nodes_from_graph_def(graph_def, fitting_net_pattern) for key in fitting_net_nodes.keys(): assert key.find('bias') > 0 or key.find('matrix') > 0 or key.find( 'idt') > 0, "currently, only support weight matrix, bias and idt at the model compression process!" @@ -277,7 +286,7 @@ def get_fitting_net_nodes(model_file : str) -> Dict: return get_fitting_net_nodes_from_graph_def(graph_def) -def get_fitting_net_variables_from_graph_def(graph_def : tf.GraphDef) -> Dict: +def get_fitting_net_variables_from_graph_def(graph_def : tf.GraphDef, suffix: str = "") -> Dict: """ Get the fitting net variables with the given tf.GraphDef object @@ -285,6 +294,8 @@ def get_fitting_net_variables_from_graph_def(graph_def : tf.GraphDef) -> Dict: ---------- graph_def The input tf.GraphDef object + suffix + suffix of the scope Returns ---------- @@ -292,7 +303,7 @@ def get_fitting_net_variables_from_graph_def(graph_def : tf.GraphDef) -> Dict: The fitting net variables within the given tf.GraphDef object """ fitting_net_variables = {} - fitting_net_nodes = get_fitting_net_nodes_from_graph_def(graph_def) + fitting_net_nodes = get_fitting_net_nodes_from_graph_def(graph_def, suffix=suffix) for item in fitting_net_nodes: node = fitting_net_nodes[item] dtype= tf.as_dtype(node.dtype).as_numpy_dtype @@ -304,7 +315,7 @@ def get_fitting_net_variables_from_graph_def(graph_def : tf.GraphDef) -> Dict: fitting_net_variables[item] = np.reshape(tensor_value, tensor_shape) return fitting_net_variables -def get_fitting_net_variables(model_file : str) -> Dict: +def get_fitting_net_variables(model_file : str, suffix: str = "") -> Dict: """ Get the fitting net variables with the given frozen model(model_file) @@ -312,6 +323,8 @@ def get_fitting_net_variables(model_file : str) -> Dict: ---------- model_file The input frozen model path + suffix + suffix of the scope Returns ---------- @@ -319,7 +332,7 @@ def get_fitting_net_variables(model_file : str) -> Dict: The fitting net variables within the given frozen model """ _, graph_def = load_graph_def(model_file) - return get_fitting_net_variables_from_graph_def(graph_def) + return get_fitting_net_variables_from_graph_def(graph_def, suffix=suffix) def get_type_embedding_net_nodes_from_graph_def(graph_def: tf.GraphDef, suffix: str = "") -> Dict: