| 49 | self.ds_column = self.eval_cfg.get('ds_column') |
| 50 | |
| 51 | def run(self): |
| 52 | filename = get_infer_output_path( |
| 53 | self.model_cfg, self.dataset_cfg, |
| 54 | osp.join(self.work_dir, 'predictions')) |
| 55 | root, ext = osp.splitext(filename) |
| 56 | partial_filename = root + '_0' + ext |
| 57 | |
| 58 | if not osp.exists(osp.realpath(filename)) and not osp.exists( |
| 59 | osp.realpath(partial_filename)): |
| 60 | print(f'{filename} not found') |
| 61 | return |
| 62 | |
| 63 | dataset = build_dataset_from_cfg(self.dataset_cfg) |
| 64 | # Postprocess dataset if necessary |
| 65 | if 'dataset_postprocessor' in self.eval_cfg: |
| 66 | |
| 67 | def postprocess(sample): |
| 68 | s = sample[self.ds_column] |
| 69 | proc = TEXT_POSTPROCESSORS.get( |
| 70 | self.eval_cfg['dataset_postprocessor']['type']) |
| 71 | sample[self.ds_column] = proc(s) |
| 72 | return sample |
| 73 | |
| 74 | dataset = dataset.map(postprocess) |
| 75 | |
| 76 | # Load predictions |
| 77 | if osp.exists(osp.realpath(filename)): |
| 78 | preds = mmengine.load(filename) |
| 79 | else: |
| 80 | filename = partial_filename |
| 81 | preds, offset = {}, 0 |
| 82 | i = 1 |
| 83 | while osp.exists(osp.realpath(filename)): |
| 84 | _preds = mmengine.load(filename) |
| 85 | filename = root + f'_{i}' + ext |
| 86 | i += 1 |
| 87 | for _o in range(len(_preds)): |
| 88 | preds[str(offset)] = _preds[str(_o)] |
| 89 | offset += 1 |
| 90 | pred_strs = [preds[str(i)]['prediction'] for i in range(len(preds))] |
| 91 | |
| 92 | # Postprocess predictions if necessary |
| 93 | if 'pred_postprocessor' in self.eval_cfg: |
| 94 | proc = TEXT_POSTPROCESSORS.get( |
| 95 | self.eval_cfg['pred_postprocessor']['type']) |
| 96 | pred_strs = [proc(s) for s in pred_strs] |
| 97 | |
| 98 | if self.ds_split: |
| 99 | references = dataset[self.ds_split][self.ds_column] |
| 100 | else: |
| 101 | references = dataset[self.ds_column] |
| 102 | |
| 103 | if len(pred_strs) != len(references): |
| 104 | print('length mismatch') |
| 105 | return |
| 106 | |
| 107 | # combine cases |
| 108 | allcase, badcase = [], [] |