From 91b79f6929a04d6e47f517b4c1aa9b7b570558fd Mon Sep 17 00:00:00 2001 From: WenmuZhou <572459439@qq.com> Date: Wed, 24 Nov 2021 09:28:34 +0000 Subject: [PATCH 1/3] pair param with key when load trained model params --- ppocr/utils/save_load.py | 24 ++++++++++++++++++++---- 1 file changed, 20 insertions(+), 4 deletions(-) diff --git a/ppocr/utils/save_load.py b/ppocr/utils/save_load.py index 702f3e9770..bf973c90eb 100644 --- a/ppocr/utils/save_load.py +++ b/ppocr/utils/save_load.py @@ -56,9 +56,25 @@ def load_model(config, model, optimizer=None): if checkpoints: if checkpoints.endswith('pdparams'): checkpoints = checkpoints.replace('.pdparams', '') - assert os.path.exists(checkpoints + ".pdopt"), \ - f"The {checkpoints}.pdopt does not exists!" - load_pretrained_params(model, checkpoints) + assert os.path.exists(checkpoints + ".pdparams"), \ + f"The {checkpoints}.pdparams does not exists!" + + # load params from trained model + params = paddle.load(checkpoints + '.pdparams') + state_dict = model.state_dict() + new_state_dict = {} + for key, value in state_dict.items(): + if key not in params: + logger.warning(f"{key} not in loaded params {params.keys()} !") + pre_value = params[key] + if list(value.shape) == list(pre_value.shape): + new_state_dict[key] = pre_value + else: + logger.warning( + f"The shape of model params {key} {value.shape} not matched with loaded params shape {pre_value.shape} !" + ) + model.set_state_dict(new_state_dict) + optim_dict = paddle.load(checkpoints + '.pdopt') if optimizer is not None: optimizer.set_state_dict(optim_dict) @@ -92,7 +108,7 @@ def load_pretrained_params(model, path): if list(state_dict[k1].shape) == list(params[k2].shape): new_state_dict[k1] = params[k2] else: - logger.info( + logger.warning( f"The shape of model params {k1} {state_dict[k1].shape} not matched with loaded params {k2} {params[k2].shape} !" ) model.set_state_dict(new_state_dict) From 4a6f7ceca6314dec02d0c0159cfddfb28c704ca8 Mon Sep 17 00:00:00 2001 From: WenmuZhou <572459439@qq.com> Date: Wed, 24 Nov 2021 09:40:33 +0000 Subject: [PATCH 2/3] pair param with key when load trained model params --- ppocr/utils/save_load.py | 19 ++++++++++--------- 1 file changed, 10 insertions(+), 9 deletions(-) diff --git a/ppocr/utils/save_load.py b/ppocr/utils/save_load.py index bf973c90eb..7bdaafd5b7 100644 --- a/ppocr/utils/save_load.py +++ b/ppocr/utils/save_load.py @@ -57,22 +57,23 @@ def load_model(config, model, optimizer=None): if checkpoints.endswith('pdparams'): checkpoints = checkpoints.replace('.pdparams', '') assert os.path.exists(checkpoints + ".pdparams"), \ - f"The {checkpoints}.pdparams does not exists!" - + "The {}.pdparams does not exists!".format(checkpoints) + # load params from trained model params = paddle.load(checkpoints + '.pdparams') state_dict = model.state_dict() new_state_dict = {} for key, value in state_dict.items(): if key not in params: - logger.warning(f"{key} not in loaded params {params.keys()} !") + logger.warning("{} not in loaded params {} !".format( + key, params.keys())) pre_value = params[key] if list(value.shape) == list(pre_value.shape): new_state_dict[key] = pre_value else: logger.warning( - f"The shape of model params {key} {value.shape} not matched with loaded params shape {pre_value.shape} !" - ) + "The shape of model params {} {} not matched with loaded params shape {} !". + format(key, value.shape, pre_value.shape)) model.set_state_dict(new_state_dict) optim_dict = paddle.load(checkpoints + '.pdopt') @@ -99,7 +100,7 @@ def load_pretrained_params(model, path): if path.endswith('pdparams'): path = path.replace('.pdparams', '') assert os.path.exists(path + ".pdparams"), \ - f"The {path}.pdparams does not exists!" + "The {}.pdparams does not exists!".format(path) params = paddle.load(path + '.pdparams') state_dict = model.state_dict() @@ -109,10 +110,10 @@ def load_pretrained_params(model, path): new_state_dict[k1] = params[k2] else: logger.warning( - f"The shape of model params {k1} {state_dict[k1].shape} not matched with loaded params {k2} {params[k2].shape} !" - ) + "The shape of model params {} {} not matched with loaded params {} {} !". + format(k1, state_dict[k1].shape, k2, params[k2].shape)) model.set_state_dict(new_state_dict) - logger.info(f"load pretrain successful from {path}") + logger.info("load pretrain successful from {}".format(path)) return model From 7f1badf74c258c116d6c6a439ea1ac162fcb97ae Mon Sep 17 00:00:00 2001 From: WenmuZhou <572459439@qq.com> Date: Wed, 24 Nov 2021 10:23:06 +0000 Subject: [PATCH 3/3] pair param with key when load trained model params --- ppocr/utils/save_load.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/ppocr/utils/save_load.py b/ppocr/utils/save_load.py index 7bdaafd5b7..4b890f6fa3 100644 --- a/ppocr/utils/save_load.py +++ b/ppocr/utils/save_load.py @@ -54,7 +54,7 @@ def load_model(config, model, optimizer=None): pretrained_model = global_config.get('pretrained_model') best_model_dict = {} if checkpoints: - if checkpoints.endswith('pdparams'): + if checkpoints.endswith('.pdparams'): checkpoints = checkpoints.replace('.pdparams', '') assert os.path.exists(checkpoints + ".pdparams"), \ "The {}.pdparams does not exists!".format(checkpoints) @@ -97,7 +97,7 @@ def load_model(config, model, optimizer=None): def load_pretrained_params(model, path): logger = get_logger() - if path.endswith('pdparams'): + if path.endswith('.pdparams'): path = path.replace('.pdparams', '') assert os.path.exists(path + ".pdparams"), \ "The {}.pdparams does not exists!".format(path)