MCPcopy Create free account
hub / github.com/Topdu/OpenOCR / main

Function main

tools/eval_rec_all_long.py:24–116  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

22
23
24def main():
25 FLAGS = parse_args()
26 cfg = Config(FLAGS.config)
27 FLAGS = vars(FLAGS)
28 opt = FLAGS.pop('opt')
29 cfg.merge_dict(FLAGS)
30 cfg.merge_dict(opt)
31
32 cfg.cfg['Global']['use_amp'] = False
33 if cfg.cfg['Global']['output_dir'][-1] == '/':
34 cfg.cfg['Global']['output_dir'] = cfg.cfg['Global']['output_dir'][:-1]
35 cfg.cfg['Global']['max_text_length'] = 200
36 cfg.cfg['Architecture']['Decoder']['max_len'] = 200
37 cfg.cfg['Metric']['name'] = 'RecMetricLong'
38 if cfg.cfg['Global']['pretrained_model'] is None:
39 cfg.cfg['Global'][
40 'pretrained_model'] = cfg.cfg['Global']['output_dir'] + '/best.pth'
41 trainer = Trainer(cfg, mode='eval')
42
43 best_model_dict = trainer.status.get('metrics', {})
44 trainer.logger.info('metric in ckpt ***************')
45 for k, v in best_model_dict.items():
46 trainer.logger.info('{}:{}'.format(k, v))
47
48 data_dirs_list = [
49 ['../ltb/long_lmdb'],
50 ]
51
52 cfg = cfg.cfg
53 file_csv = open(
54 cfg['Global']['output_dir'] + '/' +
55 cfg['Global']['output_dir'].split('/')[-1] +
56 '_result1_1_test_all_long_final_ultra_bs1.csv', 'w')
57 csv_w = csv.writer(file_csv)
58
59 for data_dirs in data_dirs_list:
60 acc_each = []
61 acc_each_num = []
62 acc_each_dis = []
63 each_long = {}
64 for datadir in data_dirs:
65 config_each = cfg.copy()
66
67 config_each['Eval']['dataset']['data_dir_list'] = [datadir]
68 valid_dataloader = build_dataloader(config_each, 'Eval',
69 trainer.logger)
70 trainer.logger.info(
71 f'{datadir} valid dataloader has {len(valid_dataloader)} iters'
72 )
73 trainer.valid_dataloader = valid_dataloader
74 metric = trainer.eval()
75 acc_each.append(metric['acc'] * 100)
76 acc_each_dis.append(metric['norm_edit_dis'])
77 acc_each_num.append(metric['all_num'])
78
79 trainer.logger.info('metric eval ***************')
80 for k, v in metric.items():
81 trainer.logger.info('{}:{}'.format(k, v))

Callers 1

Calls 7

merge_dictMethod · 0.95
evalMethod · 0.95
ConfigClass · 0.90
TrainerClass · 0.90
build_dataloaderFunction · 0.90
getMethod · 0.80
parse_argsFunction · 0.70

Tested by

no test coverage detected