Team Ai
Apppublic

mrm8488/PromptSource

sourceHugging Faceupdated 5y agoView on Hugging Face
4likes
preview_promptsource.py106 linesDownload Raw Back to seqio_tasks
1import csv2from typing import List, Optional, Tuple3 4import pkg_resources5 6# from rich import inspect7from rich.pretty import pprint8 9from promptsource.templates import TemplateCollection10 11 12def preview() -> None:13    experiment_path = pkg_resources.resource_filename(__name__, "experiment_D4.csv")14    gsheet = {}15    d4_train: List[Tuple[str, Optional[str]]] = []16    d4_eval: List[Tuple[str, Optional[str]]] = []17    d3_train_gpt: List[Tuple[str, Optional[str]]] = []18    d3_train_sglue: List[Tuple[str, Optional[str]]] = []19    experiment_path = pkg_resources.resource_filename(__name__, "experiment_D4.csv")20    with open(experiment_path) as exp_file:21        reader = csv.DictReader(exp_file)22        for row in reader:23            if row["skip"]:24                continue25            if row["subset"] == "":26                row["subset"] = None  # to match promptsource.Template object27            dataset_subset = (row["HF_name"], row["subset"])28            if row["do_train"] == "TRUE":29                d4_train.append(dataset_subset)30            if row["do_eval"] == "TRUE":31                d4_eval.append(dataset_subset)32            if row["D3_do_train"] == "TRUE" and "GPT" in row["seed_paper"]:33                d3_train_gpt.append(dataset_subset)34            if row["D3_do_train"] == "TRUE" and row["HF_name"] == "super_glue":35                d3_train_sglue.append(dataset_subset)36            gsheet[dataset_subset] = row37    all_datasets = d4_train + d4_eval + d3_train_gpt + d3_train_sglue38    print(f"Number of non-desk-rejected datasets = {len(all_datasets)}")39    print(f"Number of training sets = {len(d4_train)}")40    print(f"Number of evaluation sets = {len(d4_eval)}")41 42    template_collection = TemplateCollection()43    output = []44    missing_og_flags = []45    missing_metrics = []46    for dataset_name, subset_name in template_collection.keys:47        ds_name = (dataset_name, subset_name)48        if ds_name not in d4_eval:49            template_collection.remove(dataset_name, subset_name)50            continue51        OG = 052        non_OG = 053        dataset = template_collection.get_dataset(dataset_name, subset_name)54        for template_name in dataset.all_template_names:55            template = dataset[template_name]56            # if dataset_name == 'ropes':57            #     inspect(template.metadata)58            if not template.metadata.metrics:59                missing_metrics.append(f"{dataset_name}/{subset_name}/{template_name}")60 61            if template.metadata.original_task is True:62                OG += 163            elif template.metadata.original_task is False:64                non_OG += 165            elif template.metadata.original_task is None:66                missing_og_flags.append(dataset_name + "/" + template_name)67                continue68 69        train_size = gsheet[ds_name]["train_size"]70        if train_size == "":71            train_size = 072        else:73            train_size = int(train_size)74 75        adjusted_train_size = train_size // len(dataset.all_template_names)76 77        output.append(78            (79                f"{dataset_name} {subset_name if subset_name else ''}",80                f"{OG}-{non_OG}",81                f"{train_size:,}    {adjusted_train_size:,}",82            )83        )84 85    pprint(output)86    print(len(template_collection))87 88    print("Missing metrics:")89    pprint(missing_metrics)90 91    print("Missing original task flags:")92    pprint(missing_og_flags)93 94    # # print(d4_train_mixture)95    # print(f"Number of training templates = {len(d4_train_mixture)}")96    # # print(d4_eval_mixture)97    # print(f"Number of evaluation templates = {len(d4_eval_mixture)}")98    # # for i in seqio.TaskRegistry.names():99    # #     print(i)100    # print(f"Number of SeqIO registered templates = {len(seqio.TaskRegistry.names())}")101    # print("^ includes non-original task templates which are excluded from the eval mixture")102 103 104if __name__ == "__main__":105    preview()106