()
| 22 | |
| 23 | |
| 24 | def 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)) |
no test coverage detected