main.py 1.52 KB
Newer Older
Jason Phang's avatar
Jason Phang committed
1
2
import argparse
import json
Jason Phang's avatar
seed  
Jason Phang committed
3
4
import numpy as np
import random
Leo Gao's avatar
Leo Gao committed
5

Jason Phang's avatar
lib  
Jason Phang committed
6
7
from lm_eval import models, tasks

Leo Gao's avatar
Leo Gao committed
8

Jason Phang's avatar
Jason Phang committed
9
10
11
12
13
14
def parse_args():
    parser = argparse.ArgumentParser()
    parser.add_argument('--model', required=True)
    parser.add_argument('--model_args', default="")
    parser.add_argument('--tasks', default="all_tasks")
    parser.add_argument('--provide_description', action="store_true")
Jason Phang's avatar
lib  
Jason Phang committed
15
    parser.add_argument('--num_fewshot', type=int, default=1)
Jason Phang's avatar
seed  
Jason Phang committed
16
    parser.add_argument('--seed', type=int, default=1234)
Jason Phang's avatar
Jason Phang committed
17
    parser.add_argument('--output_path', default=None)
Jason Phang's avatar
Jason Phang committed
18
19
20
21
22
    return parser.parse_args()


def main():
    args = parse_args()
Jason Phang's avatar
seed  
Jason Phang committed
23
24
25
    random.seed(args.seed)
    np.random.seed(args.seed)

Jason Phang's avatar
lib  
Jason Phang committed
26
    lm = models.get_model(args.model).create_from_arg_string(args.model_args)
Jason Phang's avatar
Jason Phang committed
27
28
29
30
    if args.tasks == "all_tasks":
        task_names = tasks.ALL_TASKS
    else:
        task_names = args.tasks.split(",")
Jason Phang's avatar
lib  
Jason Phang committed
31
    task_dict = {
Jason Phang's avatar
Jason Phang committed
32
33
34
35
        task_name: tasks.get_task(task_name)()
        for task_name in task_names
    }
    results = {}
Jason Phang's avatar
lib  
Jason Phang committed
36
    for task_name, task in task_dict.items():
Jason Phang's avatar
Jason Phang committed
37
38
39
40
        if not task.has_validation_docs():
            continue
        result = task.evaluate(
            docs=task.validation_docs(),
Jason Phang's avatar
lib  
Jason Phang committed
41
            lm=lm,
Jason Phang's avatar
Jason Phang committed
42
            provide_description=args.provide_description,
Jason Phang's avatar
lib  
Jason Phang committed
43
            num_fewshot=args.num_fewshot,
Jason Phang's avatar
Jason Phang committed
44
45
        )
        results[task_name] = result
Jason Phang's avatar
Jason Phang committed
46
47
48
49
50
    dumped = json.dumps(results, indent=2)
    print(dumped)
    if args.output_path:
        with open(args.output_path, "w") as f:
            f.write(dumped)
Jason Phang's avatar
Jason Phang committed
51

Jason Phang's avatar
lib  
Jason Phang committed
52

Jason Phang's avatar
Jason Phang committed
53
if __name__ == "__main__":
Jason Phang's avatar
lib  
Jason Phang committed
54
    main()