mrm8488/PromptSource
4
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 