From 05dceedf328d6c7fbeb1ad7d00e83b8bbb64e37f Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Thu, 1 Sep 2022 16:43:35 -0400 Subject: [PATCH 1/3] support initilize parameters from a fitting with suffix Note: commonly there is no suffix for a fitting. --- deepmd/fit/ener.py | 2 +- deepmd/utils/graph.py | 21 +++++++++++++++------ 2 files changed, 16 insertions(+), 7 deletions(-) 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/utils/graph.py b/deepmd/utils/graph.py index fafca75f20..394a92480e 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 @@ -252,7 +252,14 @@ def get_fitting_net_nodes_from_graph_def(graph_def: tf.GraphDef) -> Dict: 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 +284,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 +292,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 +301,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 +313,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) @@ -319,7 +328,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: From 130f24bd06c7d77bbb7128d886220ce65da6d309 Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Thu, 1 Sep 2022 16:48:18 -0400 Subject: [PATCH 2/3] update docstrings Signed-off-by: Jinzhe Zeng --- deepmd/utils/graph.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/deepmd/utils/graph.py b/deepmd/utils/graph.py index 394a92480e..10e13730c5 100644 --- a/deepmd/utils/graph.py +++ b/deepmd/utils/graph.py @@ -246,6 +246,8 @@ def get_fitting_net_nodes_from_graph_def(graph_def: tf.GraphDef, suffix: str = " ---------- graph_def The input tf.GraphDef object + suffix + suffix of the scope Returns ---------- @@ -321,6 +323,8 @@ def get_fitting_net_variables(model_file : str, suffix: str = "") -> Dict: ---------- model_file The input frozen model path + suffix + suffix of the scope Returns ---------- From e2493f00d70765a381da18de732e0688300d0927 Mon Sep 17 00:00:00 2001 From: Jinzhe Zeng Date: Thu, 1 Sep 2022 16:50:15 -0400 Subject: [PATCH 3/3] add support for dipole and polar --- deepmd/fit/dipole.py | 4 ++-- deepmd/fit/polar.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) 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/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: