Fix MMLU benchmark to auto-download data and resolve path issue (#18486)
This commit is contained in:
@@ -1,6 +1,8 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import json
|
import json
|
||||||
import os
|
import os
|
||||||
|
import subprocess
|
||||||
|
import tarfile
|
||||||
import time
|
import time
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -13,6 +15,8 @@ from sglang.test.test_utils import (
|
|||||||
select_sglang_backend,
|
select_sglang_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
choices = ["A", "B", "C", "D"]
|
choices = ["A", "B", "C", "D"]
|
||||||
|
|
||||||
tokenizer = tiktoken.encoding_for_model("gpt-3.5-turbo")
|
tokenizer = tiktoken.encoding_for_model("gpt-3.5-turbo")
|
||||||
@@ -48,6 +52,28 @@ def gen_prompt(train_df, subject, k=-1):
|
|||||||
return prompt
|
return prompt
|
||||||
|
|
||||||
|
|
||||||
|
def download_data(data_dir):
|
||||||
|
"""Download and extract MMLU data if it doesn't exist."""
|
||||||
|
if os.path.isdir(os.path.join(data_dir, "test")):
|
||||||
|
return
|
||||||
|
print(f"Data not found at {data_dir}. Downloading...")
|
||||||
|
os.makedirs(data_dir, exist_ok=True)
|
||||||
|
tar_path = os.path.join(data_dir, "data.tar")
|
||||||
|
subprocess.check_call(
|
||||||
|
["wget", "-O", tar_path, "https://people.eecs.berkeley.edu/~hendrycks/data.tar"]
|
||||||
|
)
|
||||||
|
with tarfile.open(tar_path) as tar:
|
||||||
|
tar.extractall(path=data_dir, filter="data")
|
||||||
|
# The tarball extracts into a "data/" subdirectory; move contents up if needed
|
||||||
|
nested = os.path.join(data_dir, "data")
|
||||||
|
if os.path.isdir(nested):
|
||||||
|
for item in os.listdir(nested):
|
||||||
|
os.rename(os.path.join(nested, item), os.path.join(data_dir, item))
|
||||||
|
os.rmdir(nested)
|
||||||
|
os.remove(tar_path)
|
||||||
|
print("Download complete.")
|
||||||
|
|
||||||
|
|
||||||
def main(args):
|
def main(args):
|
||||||
subjects = sorted(
|
subjects = sorted(
|
||||||
[
|
[
|
||||||
@@ -174,8 +200,11 @@ def main(args):
|
|||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
parser = argparse.ArgumentParser()
|
parser = argparse.ArgumentParser()
|
||||||
parser.add_argument("--ntrain", "-k", type=int, default=5)
|
parser.add_argument("--ntrain", "-k", type=int, default=5)
|
||||||
parser.add_argument("--data_dir", "-d", type=str, default="data")
|
parser.add_argument(
|
||||||
|
"--data_dir", "-d", type=str, default=os.path.join(SCRIPT_DIR, "data")
|
||||||
|
)
|
||||||
parser.add_argument("--save_dir", "-s", type=str, default="results")
|
parser.add_argument("--save_dir", "-s", type=str, default="results")
|
||||||
parser.add_argument("--nsub", type=int, default=60)
|
parser.add_argument("--nsub", type=int, default=60)
|
||||||
args = add_common_sglang_args_and_parse(parser)
|
args = add_common_sglang_args_and_parse(parser)
|
||||||
|
download_data(args.data_dir)
|
||||||
main(args)
|
main(args)
|
||||||
|
|||||||
Reference in New Issue
Block a user