Team Ai
Apppublic

mrm8488/PromptSource

sourceHugging Faceupdated 5y agoView on Hugging Face
4likes
tasks.py422 linesDownload Raw Back to seqio_tasks
1import csv2import functools3from typing import Dict, List, Optional, Tuple4 5import datasets6import pkg_resources7import seqio8import t59import tensorflow as tf10from t5.data.glue_utils import get_glue_metric, get_super_glue_metric11from t5.evaluation import metrics as mt12 13import promptsource.templates14from promptsource.seqio_tasks import utils15 16 17GET_METRICS = {18    "BLEU": mt.bleu,19    "ROUGE": mt.rouge,20    "Span Squad": mt.span_squad,21    "Squad": mt.squad,22    "Trivia QA": mt.trivia_qa,23    "Accuracy": mt.accuracy,24    "Sequence Accuracy": mt.sequence_accuracy,25    "Pearson Correlation": mt.pearson_corrcoef,26    "Spearman Correlation": mt.spearman_corrcoef,27    "MultiRC": mt.multirc_f1_over_all_answers,28    "AUC": mt.auc,29    "COQA F1": mt.coqa_f1,30    "Edit Distance": mt.edit_distance,31    # "Mean Reciprocal Rank": mt.accuracy,  # NOTE not in T5?32    "Other": mt.accuracy,33    # Missing support for mean_multiclass_f1 etc. which need a num_classes parameter34}35 36MAX_EXAMPLES_PER_DATASET = 500_00037 38 39def strip_whitespace(output_or_target, example=None, is_target=False):40    """Cached tasks from promptsource all have a leading space on the ground-truth targets."""41    return output_or_target.strip()42 43 44def maybe_get_class_id_postprocessor(template):45    if template.get_fixed_answer_choices_list():46 47        def postprocess_fn(output_or_target, example=None, is_target=False):48            output_or_target = strip_whitespace(output_or_target)49            return t5.data.postprocessors.string_label_to_class_id(50                output_or_target, label_classes=template.get_fixed_answer_choices_list()51            )52 53        return postprocess_fn54 55    else:56        return strip_whitespace57 58 59def get_tf_dataset(split, shuffle_files, seed, dataset_name, subset_name, template, split_mapping):60    # HF datasets does not support file-level shuffling61    del shuffle_files, seed62    dataset = datasets.load_dataset(dataset_name, subset_name)63    dataset = dataset[split_mapping[split]]64    dataset = utils.apply_template(dataset, template)65    return utils.hf_dataset_to_tf_dataset(dataset)66 67 68def add_task(dataset_name, subset_name, template_name, task_name=None, split_mapping=None):69    template = all_templates.get_dataset(dataset_name, subset_name)[template_name]70    task_name = task_name or utils.get_task_name(dataset_name, subset_name, template_name)71 72    if dataset_name == "glue":73        metrics = get_glue_metric(subset_name)74    elif dataset_name == "super_glue":75        if subset_name in ("wsc.fixed", "multirc"):76            # TODO: WSC and MultiRC need special pre/postprocesing77            metrics = [mt.accuracy]78        else:79            metrics = get_super_glue_metric(subset_name)80    else:81        # TODO what if metric is null?82        metrics = [GET_METRICS[m] for m in template.metadata.metrics]83 84    dataset_splits = utils.get_dataset_splits(dataset_name, subset_name)85    split_mapping = split_mapping or {k: k for k in dataset_splits.keys()}86 87    dataset_fn = functools.partial(88        get_tf_dataset,89        seed=None,90        dataset_name=dataset_name,91        subset_name=subset_name,92        template=template,93        split_mapping=split_mapping,94    )95    data_source = seqio.FunctionDataSource(96        dataset_fn,97        splits=list(split_mapping.keys()),98        num_input_examples={s: dataset_splits[split_mapping[s]].num_examples for s in split_mapping.keys()},99    )100    output_features = {101        "inputs": seqio.Feature(t5.data.get_default_vocabulary(), add_eos=False, dtype=tf.int32),102        "targets": seqio.Feature(t5.data.get_default_vocabulary(), add_eos=True, dtype=tf.int32),103    }104    preprocessors = [105        seqio.preprocessors.tokenize,106        seqio.preprocessors.append_eos,107        seqio.CacheDatasetPlaceholder(required=False),108    ]109 110    # Add train and normal eval tasks111    seqio.TaskRegistry.add(112        task_name,113        data_source,114        preprocessors=preprocessors,115        output_features=output_features,116        metric_fns=metrics,117        postprocess_fn=maybe_get_class_id_postprocessor(template),118    )119 120    # Add rank classification eval task121    if template.answer_choices:122        rank_classification_preprocessor = functools.partial(123            t5.data.preprocessors.rank_classification,124            inputs_fn=lambda ex: tf.fill((len(ex["answer_choices"]),), ex["inputs"]),125            targets_fn=lambda ex: ex["answer_choices"],126            is_correct_fn=lambda ex: tf.equal(ex["answer_choices"], tf.strings.strip(ex["targets"])),127            weight_fn=lambda ex: 1.0,128        )129 130        fixed_choices = template.get_fixed_answer_choices_list()131        num_classes = len(fixed_choices) if fixed_choices else None132        seqio.TaskRegistry.add(133            task_name + "_score_eval",134            data_source,135            preprocessors=[rank_classification_preprocessor] + preprocessors,136            output_features=output_features,137            metric_fns=[functools.partial(t5.evaluation.metrics.rank_classification, num_classes=num_classes)],138            postprocess_fn=t5.data.postprocessors.rank_classification,139        )140 141 142datatset_subset_tuple = Tuple[str, Optional[str]]143d4_train: List[datatset_subset_tuple] = []144d4_eval: List[datatset_subset_tuple] = []145d3_train_gpt: List[datatset_subset_tuple] = []146d3_train_sglue: List[datatset_subset_tuple] = []147bias_fairness_eval: List[datatset_subset_tuple] = []148gsheet: Dict[datatset_subset_tuple, Dict] = {}149experiment_path = pkg_resources.resource_filename(__name__, "experiment_D4.csv")150with open(experiment_path) as exp_file:151    reader = csv.DictReader(exp_file)152    for row in reader:153        if row["skip"]:154            continue155        if row["subset"] == "":156            row["subset"] = None  # to match promptsource.Template object157        dataset_subset = (row["HF_name"], row["subset"])158        if row["do_train"] == "TRUE":159            d4_train.append(dataset_subset)160        if row["do_eval"] == "TRUE":161            d4_eval.append(dataset_subset)162        if row["D3_do_train"] == "TRUE" and "GPT" in row["seed_paper"]:163            d3_train_gpt.append(dataset_subset)164        if row["D3_do_train"] == "TRUE" and row["HF_name"] == "super_glue":165            d3_train_sglue.append(dataset_subset)166        if (167            row["do_eval"] == "TRUE"168            and row["task_by_convention"] == "bias_and_fairness"169            and row["HF_name"] != "winogender"170        ):171            bias_fairness_eval.append(dataset_subset)172        gsheet[dataset_subset] = row173all_datasets = d4_train + d4_eval + d3_train_gpt + d3_train_sglue + bias_fairness_eval174 175all_templates = promptsource.templates.TemplateCollection()176all_templates.remove("anli")  # Need to special-case ANLI due to weird split conventions177 178# 3 stages of training/ablation: D4 -> GPT -> SuperGLUE179d4_train_mixture: List[str] = []  # strings are dataset_subset_template180gpt_train_mixture: List[str] = []181sglue_train_mixture: List[str] = []182d4_eval_mixture: List[str] = []183bias_fairness_eval_mixture: List[str] = []184mixture_cap: Dict[str, int] = {}185single_original_task: Dict[Tuple[str, str], str] = {}186all_original_tasks: List[str] = []187for dataset_name, subset_name in all_templates.keys:188    if (dataset_name, subset_name) not in all_datasets:189        all_templates.remove(dataset_name, subset_name)190        continue191 192    dataset = all_templates.get_dataset(dataset_name, subset_name)193    num_templates = len(dataset.all_template_names)194    train_size = gsheet[(dataset_name, subset_name)]["train_size"]195    if train_size == "":196        train_size = 0197    else:198        train_size = int(train_size)199    if train_size > MAX_EXAMPLES_PER_DATASET:200        cap = MAX_EXAMPLES_PER_DATASET // num_templates201    else:202        cap = train_size203    for template_name in dataset.all_template_names:204        add_task(dataset_name, subset_name, template_name)205 206        template = dataset[template_name]207 208        task_name = utils.get_task_name(dataset_name, subset_name, template_name)209 210        if (dataset_name, subset_name) not in single_original_task and template.metadata.original_task:211            single_original_task[(dataset_name, subset_name)] = task_name212 213        if template.metadata.original_task:214            all_original_tasks.append(task_name)215 216        if (dataset_name, subset_name) in d4_train:217            d4_train_mixture.append(task_name)218            mixture_cap[task_name] = cap219        if (dataset_name, subset_name) in d3_train_gpt:220            gpt_train_mixture.append(task_name)221            mixture_cap[task_name] = cap222        if (dataset_name, subset_name) in d3_train_sglue:223            sglue_train_mixture.append(task_name)224            mixture_cap[task_name] = cap225        if (dataset_name, subset_name) in d4_eval:226            if template.metadata.original_task:227                d4_eval_mixture.append(task_name)228            # TODO use template.metadata.answer_choices here for rank eval229        if (dataset_name, subset_name) in bias_fairness_eval:230            bias_fairness_eval_mixture.append(task_name)231 232# Special case for ANLI, which has weirdly-named splits and rounds that should be subsets233dataset_name, subset_name = ("anli", None)234dataset = all_templates.get_dataset(dataset_name, subset_name)235for anli_round in ("r1", "r2", "r3"):236    for template_name in all_templates.get_dataset(dataset_name, subset_name).all_template_names:237        task_name = utils.get_task_name(dataset_name, subset_name, template_name) + f"_{anli_round}"238        split_mapping = {239            "train": f"train_{anli_round}",240            "validation": f"dev_{anli_round}",241            "test": f"test_{anli_round}",242        }243        add_task(dataset_name, subset_name, template_name, task_name, split_mapping)244 245        template = dataset[template_name]246        if template.metadata.original_task:247            d4_eval_mixture.append(task_name)  # TODO or add to ANLI special mixture248        # TODO use template.metadata.answer_choices here for rank eval249 250 251TASK_BLACKLIST = [252    # Tasks which often tokenize to > 1024 tokens currently253    "hotpot_qa_distractor_Generate_Explanations",254    "hotpot_qa_fullwiki_Generate_Explanations",255    "hotpot_qa_distractor_Generate_Answer_and_Explanations",256    "hotpot_qa_fullwiki_Generate_Answer_and_Explanations",257    "hotpot_qa_fullwiki_Generate_Answer",258    "hotpot_qa_distractor_Generate_Answer",259    "hotpot_qa_distractor_Generate_Title_2",260    "hotpot_qa_fullwiki_Generate_Title_2",261    "hotpot_qa_fullwiki_Generate_Title_1",262    "hotpot_qa_distractor_Generate_Title_1",263    "hotpot_qa_distractor_Generate_Question",264    "hotpot_qa_fullwiki_Generate_Question",265    "tab_fact_tab_fact_tab_fact_3",266    "tab_fact_tab_fact_tab_fact_2",267    "tab_fact_tab_fact_tab_fact_1",268    "tab_fact_tab_fact_tab_fact_7",269    "tab_fact_tab_fact_tab_fact_4",270    "tab_fact_tab_fact_tab_fact_5",271    "tab_fact_tab_fact_tab_fact_6",272    "wiki_hop_masked_Choose_Best_Object_Candidate",273    "wiki_hop_masked_Indirect_Question_about_Birthplace_Citizenship_Place_of_Death",274    "narrativeqa_Template_05",275    "ecthr_cases_alleged_violation_prediction_silver_rationales",276    # Tasks with broken cached files277    "gigaword_summarize_",278]279 280# Tasks that failed caching (won't try to fix them for now) - remove when we are done281D4_TRAIN_SCORE_EVAL_TASK_BLACKLIST = [282    "amazon_polarity_Is_this_product_review_positive_score_eval",283    "amazon_polarity_Is_this_review_negative_score_eval",284    "amazon_polarity_Is_this_review_score_eval",285    "amazon_polarity_User_recommend_this_product_score_eval",286    "amazon_polarity_convey_negative_or_positive_sentiment_score_eval",287    "amazon_polarity_flattering_or_not_score_eval",288    "amazon_polarity_negative_or_positive_tone_score_eval",289    "amazon_polarity_user_satisfied_score_eval",290    "amazon_polarity_would_you_buy_score_eval",291    "dbpedia_14_given_a_choice_of_categories__score_eval",292    "dbpedia_14_given_list_what_category_does_the_paragraph_belong_to_score_eval",293    "dbpedia_14_pick_one_category_for_the_following_text_score_eval",294    "wiki_hop_original_choose_best_object_affirmative_1_score_eval",295    "wiki_hop_original_choose_best_object_affirmative_2_score_eval",296    "wiki_hop_original_choose_best_object_affirmative_3_score_eval",297    "wiki_hop_original_choose_best_object_interrogative_1_score_eval",298    "wiki_hop_original_choose_best_object_interrogative_2_score_eval",299]300 301seqio.MixtureRegistry.add(302    "d4_train",303    [task for task in d4_train_mixture if task not in TASK_BLACKLIST],304    default_rate=lambda t: mixture_cap[t.name],305)306 307seqio.MixtureRegistry.add(308    "gpt_train",309    [task for task in gpt_train_mixture if task not in TASK_BLACKLIST],310    default_rate=lambda t: mixture_cap[t.name],311)312 313seqio.MixtureRegistry.add(314    "sglue_train",315    [task for task in sglue_train_mixture if task not in TASK_BLACKLIST],316    default_rate=lambda t: mixture_cap[t.name],317)318 319seqio.MixtureRegistry.add(320    "d4_gpt_train",321    [task for task in d4_train_mixture + gpt_train_mixture if task not in TASK_BLACKLIST],322    default_rate=lambda t: mixture_cap[t.name],323)324 325seqio.MixtureRegistry.add(326    "d4_gpt_sglue_train",327    [task for task in d4_train_mixture + gpt_train_mixture + sglue_train_mixture if task not in TASK_BLACKLIST],328    default_rate=lambda t: mixture_cap[t.name],329)330 331seqio.MixtureRegistry.add(332    "d4_eval",333    [task for task in d4_eval_mixture if task not in TASK_BLACKLIST],334    default_rate=functools.partial(seqio.mixing_rate_num_examples, maximum=500_000),335)  # eval mixture does not need to be capped336 337 338seqio.MixtureRegistry.add(339    "d4_score_eval",340    [341        task342        for task in seqio.TaskRegistry.names()343        if task.endswith("_score_eval")344        and task.split("_score_eval")[0] in d4_eval_mixture345        and task.split("_score_eval")[0] not in TASK_BLACKLIST346    ],347    default_rate=functools.partial(seqio.mixing_rate_num_examples, maximum=500_000),348)349 350# Train tasks we don't care about evaluating on351D4_TRAIN_SKIP_EVAL = [352    "paws_labeled_final",353    "adversarial_qa_dbidaf",354    "adversarial_qa_dbert",355    "duorc_ParaphraseRC",356    "dream",357    "amazon_polarity",358    "app_reviews",359    "imdb",360    "wiki_bio",361    "gigaword",362    "multi_news",363    "samsum",364    "dbpedia_14",365    "trec",366]367 368seqio.MixtureRegistry.add(369    "d4_train_eval",370    [371        task372        for task in d4_train_mixture373        if task not in TASK_BLACKLIST374        and not any([skip in task for skip in D4_TRAIN_SKIP_EVAL])375        and task in all_original_tasks376    ],377    default_rate=lambda t: mixture_cap[t.name],378)379 380seqio.MixtureRegistry.add(381    "d4_train_score_eval",382    [383        task384        for task in seqio.TaskRegistry.names()385        if task.endswith("_score_eval")386        and task.split("_score_eval")[0] in d4_train_mixture387        and task.split("_score_eval")[0] not in TASK_BLACKLIST388        and task not in D4_TRAIN_SCORE_EVAL_TASK_BLACKLIST389        and not any([skip in task for skip in D4_TRAIN_SKIP_EVAL])390        and task.split("_score_eval")[0] in all_original_tasks391    ],392    default_rate=functools.partial(seqio.mixing_rate_num_examples, maximum=500_000),393)394 395seqio.MixtureRegistry.add(396    "d4_train_one_og_prompt",397    [task for task in single_original_task.values() if task in d4_train_mixture and task not in TASK_BLACKLIST],398    default_rate=lambda t: mixture_cap[t.name],399)400 401seqio.MixtureRegistry.add(402    "d4_train_all_og_prompts",403    [task for task in all_original_tasks if task in d4_train_mixture and task not in TASK_BLACKLIST],404    default_rate=lambda t: mixture_cap[t.name],405)406 407seqio.MixtureRegistry.add(408    "bias_fairness_eval",409    bias_fairness_eval_mixture,410    default_rate=functools.partial(seqio.mixing_rate_num_examples, maximum=500_000),411)412 413seqio.MixtureRegistry.add(414    "bias_fairness_eval_score_eval",415    [416        task417        for task in seqio.TaskRegistry.names()418        if task.endswith("_score_eval") and task.split("_score_eval")[0] in bias_fairness_eval_mixture419    ],420    default_rate=functools.partial(seqio.mixing_rate_num_examples, maximum=500_000),421)422