From 49958dca6176d2938d19e2a9c1196d8eb73621e5 Mon Sep 17 00:00:00 2001 From: WenmuZhou Date: Mon, 9 Nov 2020 13:27:31 +0800 Subject: [PATCH 1/4] =?UTF-8?q?=E9=80=82=E9=85=8Drc=E7=89=88=E6=9C=AC?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ppocr/utils/save_load.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/ppocr/utils/save_load.py b/ppocr/utils/save_load.py index c6d2065128..e74d8faa6f 100644 --- a/ppocr/utils/save_load.py +++ b/ppocr/utils/save_load.py @@ -89,7 +89,8 @@ def init_model(config, model, logger, optimizer=None, lr_scheduler=None): "Given dir {}.pdparams not exist.".format(checkpoints) assert os.path.exists(checkpoints + ".pdopt"), \ "Given dir {}.pdopt not exist.".format(checkpoints) - para_dict, opti_dict = paddle.load(checkpoints) + para_dict = paddle.load(checkpoints + '.pdparams') + opti_dict = paddle.load(checkpoints + '.pdopt') model.set_dict(para_dict) if optimizer is not None: optimizer.set_state_dict(opti_dict) @@ -133,8 +134,8 @@ def save_model(net, """ _mkdir_if_not_exist(model_path, logger) model_prefix = os.path.join(model_path, prefix) - paddle.save(net.state_dict(), model_prefix) - paddle.save(optimizer.state_dict(), model_prefix) + paddle.save(net.state_dict(), model_prefix + '.pdparams') + paddle.save(optimizer.state_dict(), model_prefix + '.pdopt') # save metric and config with open(model_prefix + '.states', 'wb') as f: From 4eba6c0dce5fa23dee55a0fb9912f99d0088b615 Mon Sep 17 00:00:00 2001 From: WenmuZhou Date: Mon, 9 Nov 2020 13:28:15 +0800 Subject: [PATCH 2/4] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E4=B8=80=E4=BA=9B?= =?UTF-8?q?=E5=AF=BC=E8=87=B4=E4=B8=8D=E5=8F=AF=E7=94=A8=E7=9A=84bug?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/eval.py | 18 +++++------------- 1 file changed, 5 insertions(+), 13 deletions(-) diff --git a/tools/eval.py b/tools/eval.py index 07181ee75d..16cfe532aa 100755 --- a/tools/eval.py +++ b/tools/eval.py @@ -23,12 +23,8 @@ __dir__ = os.path.dirname(os.path.abspath(__file__)) sys.path.append(__dir__) sys.path.append(os.path.abspath(os.path.join(__dir__, '..'))) -import paddle -# paddle.manual_seed(2) - -from ppocr.utils.logging import get_logger from ppocr.data import build_dataloader -from ppocr.modeling import build_model +from ppocr.modeling.architectures import build_model from ppocr.postprocess import build_post_process from ppocr.metrics import build_metric from ppocr.utils.save_load import init_model @@ -39,8 +35,7 @@ import tools.program as program def main(): global_config = config['Global'] # build dataloader - eval_loader, _ = build_dataloader(config['EVAL'], device, False, - global_config) + valid_dataloader = build_dataloader(config, 'Eval', device, logger) # build post process post_process_class = build_post_process(config['PostProcess'], @@ -63,16 +58,13 @@ def main(): eval_class = build_metric(config['Metric']) # start eval - metirc = program.eval(model, eval_loader, post_process_class, eval_class) + metirc = program.eval(model, valid_dataloader, post_process_class, + eval_class) logger.info('metric eval ***************') for k, v in metirc.items(): logger.info('{}:{}'.format(k, v)) if __name__ == '__main__': - device, config = program.preprocess() - paddle.disable_static(device) - - logger = get_logger() - print_dict(config, logger) + config, device, logger, vdl_writer = program.preprocess() main() From 672318256cc0d1a4cb2d245950de1f368a2bc7e6 Mon Sep 17 00:00:00 2001 From: WenmuZhou Date: Mon, 9 Nov 2020 13:28:46 +0800 Subject: [PATCH 3/4] =?UTF-8?q?=E5=88=A0=E9=99=A4eval=E5=A4=9A=E4=BD=99?= =?UTF-8?q?=E7=9A=84=E5=8F=82=E6=95=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/program.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/tools/program.py b/tools/program.py index 41acb8665a..8bae0fd5d1 100755 --- a/tools/program.py +++ b/tools/program.py @@ -231,7 +231,7 @@ def train(config, if global_step > start_eval_step and \ (global_step - start_eval_step) % eval_batch_step == 0 and dist.get_rank() == 0: cur_metirc = eval(model, valid_dataloader, post_process_class, - eval_class, logger, print_batch_step) + eval_class) cur_metirc_str = 'cur metirc, {}'.format(', '.join( ['{}: {}'.format(k, v) for k, v in cur_metirc.items()])) logger.info(cur_metirc_str) @@ -293,8 +293,7 @@ def train(config, return -def eval(model, valid_dataloader, post_process_class, eval_class, logger, - print_batch_step): +def eval(model, valid_dataloader, post_process_class, eval_class): model.eval() with paddle.no_grad(): total_frame = 0.0 @@ -315,9 +314,6 @@ def eval(model, valid_dataloader, post_process_class, eval_class, logger, eval_class(post_result, batch) pbar.update(1) total_frame += len(images) - # if idx % print_batch_step == 0 and dist.get_rank() == 0: - # logger.info('tackling images for eval: {}/{}'.format( - # idx, len(valid_dataloader))) # Get final metirc,eg. acc or hmean metirc = eval_class.get_metric() From b28ea0a9298caef80744d6879ff7f82abb5d0a1f Mon Sep 17 00:00:00 2001 From: WenmuZhou Date: Mon, 9 Nov 2020 13:29:09 +0800 Subject: [PATCH 4/4] =?UTF-8?q?=E6=B7=BB=E5=8A=A0opencv=E4=BE=9D=E8=B5=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- requirements.txt | 1 + 1 file changed, 1 insertion(+) diff --git a/requirements.txt b/requirements.txt index 76305d0dc6..1321896349 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,6 +2,7 @@ shapely imgaug pyclipper lmdb +opencv-python==4.2.0.32 tqdm numpy visualdl