RISE / get_best_tex.py
mrazhou's picture
Upload get_best_tex.py with huggingface_hub
91dd2ed verified
Raw
History Blame Contribute Delete
8.38 kB
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)