| import os |
| import json |
| import argparse |
| from typing import List |
| from tabulate import tabulate |
|
|
| from abc import ABC, abstractmethod |
| class Result(ABC): |
| |
| def __init__(self, res_path, com_path_len, end_flag_len, metric_key:List[str]=[]) -> None: |
| super().__init__() |
| self.ab_path = res_path |
| self.rl_path = res_path[com_path_len:-end_flag_len] |
| self.metric_key = metric_key |
|
|
| @abstractmethod |
| def read_metric(self): |
| pass |
| |
| @abstractmethod |
| def process_metric(self, metric): |
| pass |
| |
| def get(self): |
| |
| metric = self.read_metric() |
| processed_metric = self.process_metric(metric) |
|
|
| |
| res = { |
| 'ab_path': self.ab_path, |
| 'rl_path': self.rl_path, |
| 'metric': processed_metric |
| } |
| return res |
|
|
| def get_all_res_path(paths, flag): |
| all_files = [] |
| for path in paths: |
| for root, dirs, files in os.walk(path): |
| for file in files: |
| if file.endswith(flag): |
| all_files.append(os.path.join(root, file)) |
| return all_files |
|
|
| def get_len_tuple(paths:List[str], end_flag:str) -> tuple: |
| if len(paths) > 1: |
| common_path_len = len(os.path.commonpath(paths)) |
| end_flag_len = len(end_flag) + 1 |
| max_path_len = max(len(path) for path in paths) - common_path_len - end_flag_len |
| elif len(paths) == 1: |
| path = paths[0] |
| p_name_len = len(path.split('/')[-2]) |
| end_flag_len = len(end_flag) + 1 |
| common_path_len = len(path) - p_name_len - end_flag_len |
| max_path_len = p_name_len |
| else: |
| print('No path found! Maybe you should specify the path for [baseline] by --base !') |
| exit(-1) |
| return [common_path_len, end_flag_len, max_path_len] |
|
|
|
|
| class AllResult: |
| def __init__( |
| self, |
| paths:List[str], |
| baseline_paths:List[str], |
| end_flag:str, |
| res_cls:Result, |
| incl:List[str] = None, |
| excl:List[str] = None, |
| metric_key:List[str] = [] |
| ) -> None: |
| super().__init__() |
| self.incl = incl |
| self.excl = excl |
| self.res_cls = res_cls |
| |
| self.all_res = [] |
| |
| self.base_paths, self.base_len_tuple = self.get_base_paths(baseline_paths, end_flag) |
| self.res_paths, self.len_tuple = self.get_res_paths(paths, end_flag) |
| |
| if self.base_len_tuple[2] + 2 > self.len_tuple[2]: |
| self.len_tuple[2] = self.base_len_tuple[2] + 2 |
| |
| self.metric_key = metric_key |
|
|
| def filter_paths(self, paths:List[str]): |
| if self.incl: |
| paths = [path |
| for path in paths |
| if any(tag in path for tag in self.incl)] |
| if self.excl: |
| paths = [path |
| for path in paths |
| if all(tag not in path for tag in self.excl)] |
| if len(paths) == 0: |
| print('No result found with the given tags (incl/excl)') |
| exit(-1) |
| return paths |
| |
| def get_res_paths(self, paths, end_flag): |
| all_paths = get_all_res_path(paths, end_flag) |
| res_paths = self.filter_paths(all_paths) |
| res_paths = [path |
| for path in res_paths |
| if path not in self.base_paths] |
| len_tuple = get_len_tuple(res_paths, end_flag) |
| return res_paths, len_tuple |
| |
| def get_base_paths(self, paths, end_flag): |
| all_paths = get_all_res_path(paths, end_flag) |
| len_tuple = get_len_tuple(all_paths, end_flag) |
| return all_paths, len_tuple |
| |
| def color_best(self): |
| all_res = self.get() |
| max_info = { |
| key: {'idx': 0, 'value': 0, 'len': len(key)} |
| for key in all_res[0]['metric'].keys() |
| } |
| for i, res in enumerate(all_res): |
| for key, value in res['metric'].items(): |
| key_len = max_info[key]['len'] |
| self.all_res[i]['metric'][key] = f'{value:^{key_len}.1f}' |
| if value > max_info[key]['value']: |
| max_info[key]['value'] = value |
| max_info[key]['idx'] = i |
| |
| for key, item in max_info.items(): |
| idx = item['idx'] |
| value = item['value'] |
| key_len = max_info[key]['len'] |
| self.all_res[idx]['metric'][key] = f'\033m\033[31m{value:^{key_len}.2f}\033[0m' |
| return max_info |
|
|
| def sorted_by(self, key:str): |
| all_res = self.get() |
| self.all_res = sorted( |
| all_res, |
| key=lambda x: x['metric'][key], |
| reverse=True |
| ) |
| |
| def header(self, all_res, delimiter='\t'): |
| if not all_res: |
| return 'name\tmetric' |
| first = '{:^{}}'.format('name', self.len_tuple[2]) |
| metric_key = self.metric_key or all_res[0]['metric'].keys() |
| others = f'{delimiter}'.join(metric_key) |
| return f'{first} | {others}' |
| |
| def print(self, delimiter='\t'): |
| all_res = self.get() |
| print(self.header(all_res, delimiter=delimiter)) |
| for res in all_res: |
| rl_path = f'{res["rl_path"]:{self.len_tuple[2]}}' |
| if self.metric_key: |
| metric = f'{delimiter}'.join([f'{res["metric"][k]}' for k in self.metric_key]) |
| else: |
| metric = f'{delimiter}'.join([f'{v}' for v in res['metric'].values()]) |
| print(f'{rl_path} | {metric}') |
| |
| def print_table(self): |
| all_res = self.get() |
| data = [] |
| headers = ['name', *self.metric_key] |
| for res in all_res: |
| data.append([res['rl_path'], *res['metric'].values()]) |
| print(tabulate(data, headers=headers, tablefmt='plain', floatfmt='.2f', numalign='center')) |
|
|
| def get(self): |
| if self.all_res: |
| return self.all_res |
|
|
| for res_path in self.res_paths: |
| res = self.res_cls( |
| res_path, |
| self.len_tuple[0], |
| self.len_tuple[1], |
| self.metric_key |
| ).get() |
| self.all_res.append(res) |
| |
| |
| for i, res_path in enumerate(self.base_paths): |
| res = self.res_cls( |
| res_path, |
| self.base_len_tuple[0], |
| self.base_len_tuple[1], |
| self.metric_key |
| ).get() |
| res['rl_path'] = f'[{res["rl_path"]}]' |
| self.all_res.append(res) |
| |
| return self.all_res |
|
|
| class CapRes(Result): |
| sub_metric_key = 'CIDEr' |
| |
| def read_metric(self): |
| with open(self.ab_path, 'r') as infile: |
| metrics = json.load(infile) |
| return metrics |
|
|
| def process_metric(self, metric): |
| processed_metric = {} |
| for key, value in metric.items(): |
| processed_metric[key] = round(value[CapRes.sub_metric_key] * 100, 2) |
| for key in self.metric_key: |
| if key not in processed_metric: |
| processed_metric[key] = 0 |
| return processed_metric |
|
|
| def main(args): |
| CapRes.sub_metric_key = args.sub_metric_key |
| all_res = AllResult( |
| args.paths, |
| args.base, |
| args.end_flag, |
| CapRes, |
| incl=args.incl, |
| excl=args.excl, |
| metric_key=['coco', 'indomain', 'neardomain', 'outdomain', 'overall', 'msrvtt', 'spotdiff', 'flickr30k'], |
| ) |
| all_res.sorted_by(args.sort_by) |
| all_res.color_best() |
| all_res.print(delimiter='& ') |
|
|
| if __name__ == '__main__': |
| parser = argparse.ArgumentParser() |
| parser.add_argument('paths_', nargs='+', default=['./checkpoints/reproduce'], help='path of result files') |
| parser.add_argument('--paths', nargs='+', default=[], help='path of result files') |
| parser.add_argument('--base', nargs='+', default=['checkpoints/main/1.0-msrvtt/'], help='path of baseline files') |
| parser.add_argument('--end_flag', type=str, default='eval.json') |
| parser.add_argument('-e', '--excl', type=str, nargs='*', default=[]) |
| parser.add_argument('-i', '--incl', type=str, nargs='*', default=[]) |
| parser.add_argument('-s', '--sort_by', type=str, default='coco') |
| parser.add_argument('-m', '--sub_metric_key', type=str, default='CIDEr') |
| args = parser.parse_args() |
| args.paths += args.paths_ |
| main(args) |
|
|