File size: 8,871 Bytes
cab74fb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 | import argparse
import torch
from almeval.datasets import build_dataset
from almeval.models import build_model
from tqdm import tqdm
from almeval.utils import *
from loguru import logger
import sys
import os
from datetime import datetime
import warnings
import time
warnings.filterwarnings("ignore", category=FutureWarning)
warnings.filterwarnings("ignore", category=DeprecationWarning)
warnings.filterwarnings("ignore", category=UserWarning)
def setup_logging(rank, model, data, work_dir):
# Create log directory
log_dir = os.path.join(work_dir, model, data, 'logs')
os.makedirs(log_dir, exist_ok=True)
log_file = os.path.join(log_dir, f'rank{rank}.log')
# Remove default console output
logger.remove()
# Add file output
log_file_handle = open(log_file, 'w', encoding='utf-8')
logger.add(log_file_handle,
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {message}",
level="INFO")
# Print rank=0 output to console as well
if rank == 0:
logger.add(sys.stdout,
format="{time:YYYY-MM-DD HH:mm:ss} | {level} | {message}",
level="INFO")
return log_file
def merge_one_dataset(args, dataset, result_file, eval_file):
model_data_dir = osp.join(args.work_dir, args.model, dataset.DATASET_NAME)
os.makedirs(model_data_dir, exist_ok=True)
if args.reeval:
perf = dataset.evaluate(result_file, method=args.eval_method)
with open(eval_file, 'w') as f:
json.dump(perf, f, indent=4)
return
tmp_files = [osp.join(model_data_dir, f'{rank}_{args.world_size}_{dataset.DATASET_NAME}.pkl') for rank in range(args.world_size)]
# Merge if all tmp_files exist
if all(osp.exists(tmpfile) for tmpfile in tmp_files):
data_all = {}
for tmpfile in tmp_files:
data_all.update(load(tmpfile))
raw_data = dataset.data
for x in raw_data:
idx = int(x['index'])
if idx not in data_all:
logger.warning(f'index {idx} not found in data_all, details: {x}')
x['prediction'] = 'null'
x['real_prompt'] = ''
continue
x['prediction'] = str(data_all[idx]['prediction'])
x['real_prompt'] = str(data_all[idx]['prompt'])
dump(raw_data, result_file)
for tmpfile in tmp_files:
os.remove(tmpfile)
logger.info(f'model {args.model}, data {dataset.DATASET_NAME}, all {args.world_size} result merged to {result_file}.')
if args.skip_eval:
logger.info(f'skip eval for {dataset.DATASET_NAME}')
return
perf = dataset.evaluate(result_file, method=args.eval_method)
with open(eval_file, 'w') as f:
json.dump(perf, f, indent=4)
logger.info(f'model {args.model}, data {dataset.DATASET_NAME} evaluated.')
def do_reeval(dataset_name, result_file='auto', method='default'):
datasets = []
for dataset_name in args.data:
d = build_dataset(dataset_name)
if isinstance(d, list):
datasets.extend(d)
else:
datasets.append(d)
for dataset in datasets:
if result_file == 'auto':
benchmark_dir = osp.join(args.work_dir, args.model, dataset.DATASET_NAME)
pred_result_file = osp.join(benchmark_dir, f'{args.model}_{dataset.DATASET_NAME}.jsonl')
else:
pred_result_file = result_file
logger.info(f'evaluating {pred_result_file} with method {method}')
perf = dataset.evaluate(pred_result_file, method=method)
with open(pred_result_file.replace('.jsonl', f'_{method}_performance.json'), 'w') as f:
json.dump(perf, f, indent=4)
def process_dataset(args, dataset, model):
# Assign different subsets to each process
model_data_dir = osp.join(args.work_dir, args.model, dataset.DATASET_NAME)
result_file = osp.join(model_data_dir, f'{args.model}_{dataset.DATASET_NAME}.jsonl')
eval_file = osp.join(model_data_dir, f'{args.model}_{dataset.DATASET_NAME}_{args.eval_method}_performance.json')
os.makedirs(model_data_dir, exist_ok=True)
rank = int(args.rank)
if os.path.exists(result_file) and not args.force_reinfer:
if args.reeval or not os.path.exists(eval_file):
if rank==0:
logger.info(f'file {result_file} exists, reevaluating...')
merge_one_dataset(args, dataset, result_file, eval_file)
return
else:
# Exit if not reeval and result file exists
return
else:
if args.debug:
dataset.set_demo_mode()
sample_indices = [i for i in range(len(dataset))]
# Distribute data to each rank
world_size = int(args.world_size)
rank = int(args.rank)
sample_indices_sub = sample_indices[rank::world_size]
tmpl = osp.join(model_data_dir, f'{rank}_{args.world_size}_{dataset.DATASET_NAME}.pkl')
out_file = tmpl.format(rank)
res = load(out_file) if osp.exists(out_file) else {}
processed_samples = 0
for i in tqdm(sample_indices_sub, disable=args.rank != 0):
msg = dataset[i]
idx = int(msg['index'])
if not args.force_reinfer:
if idx in res:
continue
if processed_samples==0:
logger.info(f'Msg example: {msg}')
real_prompt, response = model(msg)
torch.cuda.empty_cache()
if response is None:
continue
# we need response and prompt, because model may change prompt
res[idx] = {
'prompt': real_prompt,
'prediction': response,
}
processed_samples += 1
if processed_samples % 20 == 0:
dump(res, out_file)
dump(res, out_file)
# Write a file to indicate this rank is done
with open(osp.join(model_data_dir, f'{rank}_{args.world_size}_{dataset.DATASET_NAME}.done'), 'w') as f:
f.write('done')
# Rank 0 needs to wait for other ranks to finish, then merge results
time_elapsed = 0
if rank == 0:
while True:
all_success_files = [osp.join(model_data_dir, f'{rank}_{args.world_size}_{dataset.DATASET_NAME}.done')
for rank in range(args.world_size)]
if len(all_success_files) == args.world_size and \
all(osp.exists(success_file) for success_file in all_success_files):
# Delete all done files
for success_file in all_success_files:
os.remove(success_file)
break
else:
time.sleep(10)
time_elapsed += 10
logger.info(f'waiting for other ranks to finish, time elapsed: {time_elapsed}s')
merge_one_dataset(args, dataset, result_file, eval_file)
def main(args):
datasets = []
for dataset_name in args.data:
d = build_dataset(dataset_name)
if isinstance(d, list):
datasets.extend(d)
else:
datasets.append(d)
model = build_model(args.model)
logger.info(f"Datasets: {datasets}")
for dataset in datasets:
setup_logging(args.rank, args.model, dataset.DATASET_NAME, args.work_dir)
logger.info(f"Running {args.model} on dataset: {dataset.DATASET_NAME}")
process_dataset(args, dataset, model)
if __name__ == '__main__':
parser = argparse.ArgumentParser()
parser.add_argument('--data', type=str, nargs='+', required=True, help='List of dataset names')
parser.add_argument('--model', type=str, help='Model name')
parser.add_argument('--work-dir', type=str, default='./eval_result', help='Working directory')
parser.add_argument('--rank', type=int, default=0, help='Current GPU rank')
parser.add_argument('--world-size', type=int, default=1, help='Total number of GPUs')
parser.add_argument('--reeval', action='store_true', help='Whether to re-evaluate')
parser.add_argument('--eval-file', type=str, default='auto', help='Evaluation file path, used when run_eval_only=True')
parser.add_argument('--debug', action='store_true', help='Debug mode, only run 10 samples per dataset')
parser.add_argument('--eval-method', type=str, default='default', help='Evaluation method')
parser.add_argument('--force-reinfer', action='store_true', help='Whether to force re-inference')
parser.add_argument('--skip-eval', action='store_true', help='Whether to skip evaluation')
args = parser.parse_args()
if args.reeval:
do_reeval(args.data, args.eval_file, args.eval_method)
else:
main(args)
|