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): # read and process metric = self.read_metric() processed_metric = self.process_metric(metric) # set result 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 # +1 for the '/' 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 # +1 for the '/' 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) # baseline 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)