bigscience/promptsource
105
1import argparse2import textwrap3 4from promptsource.templates import TemplateCollection, INCLUDED_USERS5from promptsource.utils import get_dataset6 7 8parser = argparse.ArgumentParser(description="Process some integers.")9parser.add_argument("dataset_path", type=str, help="path to dataset name")10 11args = parser.parse_args()12if "templates.yaml" not in args.dataset_path:13 exit()14 15path = args.dataset_path.split("/")16 17if path[2] in INCLUDED_USERS:18 print("Skipping showing templates for community dataset.")19else:20 dataset_name = path[2]21 subset_name = path[3] if len(path) == 5 else ""22 23 template_collection = TemplateCollection()24 25 dataset = get_dataset(dataset_name, subset_name)26 splits = list(dataset.keys())27 28 dataset_templates = template_collection.get_dataset(dataset_name, subset_name)29 template_list = dataset_templates.all_template_names30 31 width = 8032 print("DATASET ", args.dataset_path)33 34 # First show all the templates.35 for template_name in template_list:36 template = dataset_templates[template_name]37 print("TEMPLATE")38 print("NAME:", template_name)39 print("Is Original Task: ", template.metadata.original_task)40 print(template.jinja)41 print()42 43 # Show examples of the templates.44 for template_name in template_list:45 template = dataset_templates[template_name]46 print()47 print("TEMPLATE")48 print("NAME:", template_name)49 print("REFERENCE:", template.reference)50 print("--------")51 print()52 print(template.jinja)53 print()54 55 for split_name in splits:56 dataset_split = dataset[split_name]57 58 print_counter = 059 for example in dataset_split:60 print("\t--------")61 print("\tSplit ", split_name)62 print("\tExample ", example)63 print("\t--------")64 output = template.apply(example)65 if output[0].strip() == "" or (len(output) > 1 and output[1].strip() == ""):66 print("\t Blank result")67 continue68 69 xp, yp = output70 print()71 print("\tPrompt | X")72 for line in textwrap.wrap(xp, width=width, replace_whitespace=False):73 print("\t", line.replace("\n", "\n\t"))74 print()75 print("\tY")76 for line in textwrap.wrap(yp, width=width, replace_whitespace=False):77 print("\t", line.replace("\n", "\n\t"))78 79 print_counter += 180 if print_counter >= 10:81 break82 