bigscience/promptsource
105
1import argparse2import functools3import multiprocessing4import os5import textwrap6from hashlib import sha2567from multiprocessing import Manager, Pool8 9import pandas as pd10import plotly.express as px11import streamlit as st12from datasets import get_dataset_infos13from datasets.info import DatasetInfosDict14from pygments import highlight15from pygments.formatters import HtmlFormatter16from pygments.lexers import DjangoLexer17 18from promptsource import DEFAULT_PROMPTSOURCE_CACHE_HOME19from promptsource.session import _get_state20from promptsource.templates import INCLUDED_USERS, LANGUAGES, METRICS, DatasetTemplates, Template, TemplateCollection21from promptsource.utils import (22 get_dataset,23 get_dataset_confs,24 list_datasets,25 removeHyphen,26 renameDatasetColumn,27 render_features,28)29 30 31DATASET_INFOS_CACHE_DIR = os.path.join(DEFAULT_PROMPTSOURCE_CACHE_HOME, "DATASET_INFOS")32os.makedirs(DATASET_INFOS_CACHE_DIR, exist_ok=True)33 34# Python 3.8 switched the default start method from fork to spawn. OS X also has35# some issues related to fork, eee, e.g., https://github.com/bigscience-workshop/promptsource/issues/57236# so we make sure we always use spawn for consistency37multiprocessing.set_start_method("spawn", force=True)38 39 40def get_infos(all_infos, d_name):41 """42 Wrapper for mutliprocess-loading of dataset infos43 44 :param all_infos: multiprocess-safe dictionary45 :param d_name: dataset name46 """47 d_name_bytes = d_name.encode("utf-8")48 d_name_hash = sha256(d_name_bytes)49 foldername = os.path.join(DATASET_INFOS_CACHE_DIR, d_name_hash.hexdigest())50 if os.path.isdir(foldername):51 infos_dict = DatasetInfosDict.from_directory(foldername)52 else:53 infos = get_dataset_infos(d_name)54 infos_dict = DatasetInfosDict(infos)55 os.makedirs(foldername)56 infos_dict.write_to_directory(foldername)57 all_infos[d_name] = infos_dict58 59 60def format_language(tag):61 """62 Formats a language tag for display in the UI.63 64 For example, if the tag is "en", then the function returns "en (English)"65 :param tag: language tag66 :return: formatted language name67 """68 return tag + " (" + LANGUAGES[tag] + ")"69 70 71# add an argument for read-only72# At the moment, streamlit does not handle python script arguments gracefully.73# Thus, for read-only mode, you have to type one of the below two:74# streamlit run promptsource/app.py -- -r75# streamlit run promptsource/app.py -- --read-only76# Check https://github.com/streamlit/streamlit/issues/337 for more information.77parser = argparse.ArgumentParser(description="run app.py with args")78parser.add_argument("-r", "--read-only", action="store_true", help="whether to run it as read-only mode", default=True)79 80args = parser.parse_args()81if args.read_only:82 select_options = ["Helicopter view", "Prompted dataset viewer"]83 side_bar_title_prefix = "Promptsource (Read only)"84else:85 select_options = ["Helicopter view", "Prompted dataset viewer", "Sourcing"]86 side_bar_title_prefix = "Promptsource"87 88#89# Cache functions90#91get_dataset = st.cache(allow_output_mutation=True)(get_dataset)92get_dataset_confs = st.cache(get_dataset_confs)93list_datasets = st.cache(list_datasets)94 95 96def run_app():97 #98 # Loads session state99 #100 state = _get_state()101 102 def reset_template_state():103 state.template_name = None104 state.jinja = None105 state.reference = None106 107 #108 # Initial page setup109 #110 st.set_page_config(page_title="Promptsource", layout="wide")111 st.sidebar.markdown(112 "<center><a href='https://github.com/bigscience-workshop/promptsource'>π»Github - Promptsource\n\n</a></center>",113 unsafe_allow_html=True,114 )115 mode = st.sidebar.selectbox(116 label="Choose a mode",117 options=select_options,118 index=0,119 key="mode_select",120 )121 st.sidebar.title(f"{side_bar_title_prefix} πΈ - {mode}")122 123 #124 # Adds pygments styles to the page.125 #126 st.markdown(127 "<style>" + HtmlFormatter(style="friendly").get_style_defs(".highlight") + "</style>", unsafe_allow_html=True128 )129 130 WIDTH = 140131 132 def show_jinja(t, width=WIDTH):133 def replace_linebreaks(t):134 """135 st.write does not handle double breaklines very well. When it encounters `\n\n`, it exit the curent <div> block.136 Explicitely replacing all `\n` with their html equivalent to bypass this issue.137 Also stripping the trailing `\n` first.138 """139 return t.strip("\n").replace("\n", "<br/>")140 141 wrap = textwrap.fill(t, width=width, replace_whitespace=False)142 out = highlight(wrap, DjangoLexer(), HtmlFormatter())143 out = replace_linebreaks(out)144 st.write(out, unsafe_allow_html=True)145 146 def show_text(t, width=WIDTH, with_markdown=False):147 wrap = [textwrap.fill(subt, width=width, replace_whitespace=False) for subt in t.split("\n")]148 wrap = "\n".join(wrap)149 if with_markdown:150 st.write(wrap, unsafe_allow_html=True)151 else:152 st.text(wrap)153 154 if mode == "Helicopter view":155 st.title("High level metrics")156 st.write("This will take a minute to collect.")157 st.write(158 "If you want to contribute, please refer to the instructions in "159 + "[Contributing](https://github.com/bigscience-workshop/promptsource/blob/main/CONTRIBUTING.md)."160 )161 162 #163 # Loads template data164 #165 try:166 template_collection = TemplateCollection()167 except FileNotFoundError:168 st.error(169 "Unable to find the prompt folder!\n\n"170 "We expect the folder to be in the working directory. "171 "You might need to restart the app in the root directory of the repo."172 )173 st.stop()174 175 #176 # Global metrics177 #178 counts = template_collection.get_templates_count()179 nb_prompted_datasets = len(counts)180 st.write(f"## Number of *prompted datasets*: `{nb_prompted_datasets}`")181 nb_prompts = sum(counts.values())182 st.write(f"## Number of *prompts*: `{nb_prompts}`")183 184 #185 # Metrics per dataset/subset186 #187 # Download dataset infos (multiprocessing download)188 manager = Manager()189 all_infos = manager.dict()190 all_datasets = list(set([t[0] for t in template_collection.keys]))191 192 pool = Pool(processes=multiprocessing.cpu_count())193 pool.map(functools.partial(get_infos, all_infos), all_datasets)194 pool.close()195 pool.join()196 197 results = []198 for (dataset_name, subset_name) in template_collection.keys:199 # Collect split sizes (train, validation and test)200 if dataset_name not in all_infos:201 infos = get_dataset_infos(dataset_name)202 all_infos[dataset_name] = infos203 else:204 infos = all_infos[dataset_name]205 if infos:206 if subset_name is None:207 subset_infos = infos[list(infos.keys())[0]]208 else:209 subset_infos = infos[subset_name]210 211 try:212 split_sizes = {k: v.num_examples for k, v in subset_infos.splits.items()}213 except Exception:214 # Fixing bug in some community datasets.215 # For simplicity, just filling `split_sizes` with nothing, so the displayed split sizes will be 0.216 split_sizes = {}217 else:218 split_sizes = {}219 220 # Collect template counts, original task counts and names221 dataset_templates = template_collection.get_dataset(dataset_name, subset_name)222 results.append(223 {224 "Dataset name": dataset_name,225 "Subset name": "β
" if subset_name is None else subset_name,226 "Train size": split_sizes["train"] if "train" in split_sizes else 0,227 "Validation size": split_sizes["validation"] if "validation" in split_sizes else 0,228 "Test size": split_sizes["test"] if "test" in split_sizes else 0,229 "Number of prompts": len(dataset_templates),230 "Number of original task prompts": sum(231 [bool(t.metadata.original_task) for t in dataset_templates.templates.values()]232 ),233 "Prompt names": [t.name for t in dataset_templates.templates.values()],234 }235 )236 results_df = pd.DataFrame(results)237 results_df.sort_values(["Number of prompts"], inplace=True, ascending=False)238 results_df.reset_index(drop=True, inplace=True)239 240 nb_training_instances = results_df["Train size"].sum()241 st.write(f"## Number of *training instances*: `{nb_training_instances}`")242 243 plot_df = results_df[["Dataset name", "Subset name", "Train size", "Number of prompts"]].copy()244 plot_df["Name"] = plot_df["Dataset name"] + " - " + plot_df["Subset name"]245 plot_df.sort_values(["Train size"], inplace=True, ascending=False)246 fig = px.bar(247 plot_df,248 x="Name",249 y="Train size",250 hover_data=["Dataset name", "Subset name", "Number of prompts"],251 log_y=True,252 title="Number of training instances per data(sub)set - y-axis is in logscale",253 )254 fig.update_xaxes(visible=False, showticklabels=False)255 st.plotly_chart(fig, use_container_width=True)256 st.write(257 f"- Top 3 training subsets account for `{100 * plot_df[:3]['Train size'].sum() / nb_training_instances:.2f}%` of the training instances."258 )259 biggest_training_subset = plot_df.iloc[0]260 st.write(261 f"- Biggest training subset is *{biggest_training_subset['Name']}* with `{biggest_training_subset['Train size']}` instances"262 )263 smallest_training_subset = plot_df[plot_df["Train size"] > 0].iloc[-1]264 st.write(265 f"- Smallest training subset is *{smallest_training_subset['Name']}* with `{smallest_training_subset['Train size']}` instances"266 )267 268 st.markdown("***")269 st.write("Details per dataset")270 st.table(results_df)271 272 else:273 # Combining mode `Prompted dataset viewer` and `Sourcing` since the274 # backbone of the interfaces is the same275 assert mode in ["Prompted dataset viewer", "Sourcing"], ValueError(276 f"`mode` ({mode}) should be in `[Helicopter view, Prompted dataset viewer, Sourcing]`"277 )278 279 #280 # Loads dataset information281 #282 283 dataset_list = list_datasets()284 ag_news_index = dataset_list.index("ag_news")285 286 #287 # Select a dataset - starts with ag_news288 #289 dataset_key = st.sidebar.selectbox(290 "Dataset",291 dataset_list,292 key="dataset_select",293 index=ag_news_index,294 help="Select the dataset to work on.",295 )296 297 #298 # If a particular dataset is selected, loads dataset and template information299 #300 if dataset_key is not None:301 302 #303 # Check for subconfigurations (i.e. subsets)304 #305 configs = get_dataset_confs(dataset_key)306 conf_option = None307 if len(configs) > 0:308 conf_option = st.sidebar.selectbox("Subset", configs, index=0, format_func=lambda a: a.name)309 310 subset_name = str(conf_option.name) if conf_option else None311 try:312 dataset = get_dataset(dataset_key, subset_name)313 except OSError as e:314 st.error(315 f"Some datasets are not handled automatically by `datasets` and require users to download the "316 f"dataset manually. This applies to {dataset_key}{f'/{subset_name}' if subset_name is not None else ''}. "317 f"\n\nPlease download the raw dataset to `~/.cache/promptsource/{dataset_key}{f'/{subset_name}' if subset_name is not None else ''}`. "318 f"\n\nYou can choose another cache directory by overriding `PROMPTSOURCE_MANUAL_DATASET_DIR` environment "319 f"variable and downloading raw dataset to `$PROMPTSOURCE_MANUAL_DATASET_DIR/{dataset_key}{f'/{subset_name}' if subset_name is not None else ''}`"320 f"\n\nOriginal error:\n{str(e)}"321 )322 st.stop()323 324 splits = list(dataset.keys())325 index = 0326 if "train" in splits:327 index = splits.index("train")328 split = st.sidebar.selectbox("Split", splits, key="split_select", index=index)329 dataset = dataset[split]330 dataset = renameDatasetColumn(dataset)331 332 #333 # Loads template data334 #335 try:336 dataset_templates = DatasetTemplates(dataset_key, conf_option.name if conf_option else None)337 except FileNotFoundError:338 st.error(339 "Unable to find the prompt folder!\n\n"340 "We expect the folder to be in the working directory. "341 "You might need to restart the app in the root directory of the repo."342 )343 st.stop()344 345 template_list = dataset_templates.all_template_names346 num_templates = len(template_list)347 st.sidebar.write(348 "No of prompts created for "349 + f"`{dataset_key + (('/' + conf_option.name) if conf_option else '')}`"350 + f": **{str(num_templates)}**"351 )352 353 if mode == "Prompted dataset viewer":354 if num_templates > 0:355 template_name = st.sidebar.selectbox(356 "Prompt name",357 template_list,358 key="template_select",359 index=0,360 help="Select the prompt to visualize.",361 )362 363 step = 50364 example_index = st.sidebar.number_input(365 f"Select the example index (Size = {len(dataset)})",366 min_value=0,367 max_value=len(dataset) - step,368 value=0,369 step=step,370 key="example_index_number_input",371 help="Offset = 50.",372 )373 else: # mode = Sourcing374 st.sidebar.subheader("Select Example")375 example_index = st.sidebar.slider("Select the example index", 0, len(dataset) - 1)376 377 example = dataset[example_index]378 example = removeHyphen(example)379 380 st.sidebar.write(example)381 382 st.sidebar.subheader("Dataset Schema")383 rendered_features = render_features(dataset.features)384 st.sidebar.write(rendered_features)385 386 #387 # Display dataset information388 #389 st.header("Dataset: " + dataset_key + " " + (("/ " + conf_option.name) if conf_option else ""))390 391 # If we have a custom dataset change the source link to the hub392 split_dataset_key = dataset_key.split("/")393 possible_user = split_dataset_key[0]394 if len(split_dataset_key) > 1 and possible_user in INCLUDED_USERS:395 source_link = "https://huggingface.co/datasets/%s/blob/main/%s.py" % (396 dataset_key,397 split_dataset_key[-1],398 )399 else:400 source_link = "https://github.com/huggingface/datasets/blob/master/datasets/%s/%s.py" % (401 dataset_key,402 dataset_key,403 )404 405 st.markdown("*Homepage*: " + dataset.info.homepage + "\n\n*Dataset*: " + source_link)406 407 md = """408 %s409 """ % (410 dataset.info.description.replace("\\", "") if dataset_key else ""411 )412 st.markdown(md)413 414 #415 # Body of the app: display prompted examples in mode `Prompted dataset viewer`416 # or text boxes to create new prompts in mode `Sourcing`417 #418 if mode == "Prompted dataset viewer":419 #420 # Display template information421 #422 if num_templates > 0:423 template = dataset_templates[template_name]424 st.subheader("Prompt")425 st.markdown("##### Name")426 st.text(template.name)427 st.markdown("##### Reference")428 st.text(template.reference)429 st.markdown("##### Original Task? ")430 st.text(template.metadata.original_task)431 st.markdown("##### Choices in template? ")432 st.text(template.metadata.choices_in_prompt)433 st.markdown("##### Metrics")434 st.text(", ".join(template.metadata.metrics) if template.metadata.metrics else None)435 st.markdown("##### Prompt Languages")436 if template.metadata.languages:437 st.text(", ".join([format_language(tag) for tag in template.metadata.languages]))438 else:439 st.text(None)440 st.markdown("##### Answer Choices")441 if template.get_answer_choices_expr() is not None:442 show_jinja(template.get_answer_choices_expr())443 else:444 st.text(None)445 st.markdown("##### Jinja template")446 splitted_template = template.jinja.split("|||")447 st.markdown("###### Input template")448 show_jinja(splitted_template[0].strip())449 if len(splitted_template) > 1:450 st.markdown("###### Target template")451 show_jinja(splitted_template[1].strip())452 st.markdown("***")453 454 #455 # Display a couple (steps) examples456 #457 for ex_idx in range(example_index, example_index + step):458 if ex_idx >= len(dataset):459 continue460 example = dataset[ex_idx]461 example = removeHyphen(example)462 col1, _, col2 = st.beta_columns([12, 1, 12])463 with col1:464 st.write(example)465 if num_templates > 0:466 with col2:467 prompt = template.apply(example, highlight_variables=False)468 if prompt == [""]:469 st.write("β
β
β
*Blank result*")470 else:471 st.write("Input")472 show_text(prompt[0])473 if len(prompt) > 1:474 st.write("Target")475 show_text(prompt[1])476 st.markdown("***")477 else: # mode = Sourcing478 st.markdown("## Prompt Creator")479 480 #481 # Create a new template or select an existing one482 #483 col1a, col1b, _, col2 = st.beta_columns([9, 9, 1, 6])484 485 # current_templates_key and state.templates_key are keys for the templates object486 current_templates_key = (dataset_key, conf_option.name if conf_option else None)487 488 # Resets state if there has been a change in templates_key489 if state.templates_key != current_templates_key:490 state.templates_key = current_templates_key491 reset_template_state()492 493 with col1a, st.form("new_template_form"):494 new_template_name = st.text_input(495 "Create a New Prompt",496 key="new_template",497 value="",498 help="Enter name and hit enter to create a new prompt.",499 )500 new_template_submitted = st.form_submit_button("Create")501 if new_template_submitted:502 if new_template_name in dataset_templates.all_template_names:503 st.error(504 f"A prompt with the name {new_template_name} already exists "505 f"for dataset {state.templates_key}."506 )507 elif new_template_name == "":508 st.error("Need to provide a prompt name.")509 else:510 template = Template(new_template_name, "", "")511 dataset_templates.add_template(template)512 reset_template_state()513 state.template_name = new_template_name514 else:515 state.new_template_name = None516 517 with col1b, st.beta_expander("or Select Prompt", expanded=True):518 template_list = dataset_templates.all_template_names519 if state.template_name:520 index = template_list.index(state.template_name)521 else:522 index = 0523 state.template_name = st.selectbox(524 "", template_list, key="template_select", index=index, help="Select the prompt to work on."525 )526 527 if st.button("Delete Prompt", key="delete_prompt"):528 dataset_templates.remove_template(state.template_name)529 reset_template_state()530 531 variety_guideline = """532 :heavy_exclamation_mark::question:Creating a diverse set of prompts whose differences go beyond surface wordings (i.e. marginally changing 2 or 3 words) is highly encouraged.533 Ultimately, the hope is that exposing the model to such a diversity will have a non-trivial impact on the model's robustness to the prompt formulation.534 \r**To get various prompts, you can try moving the cursor along theses axes**:535 \n- **Interrogative vs affirmative form**: Ask a question about an attribute of the inputs or tell the model to decide something about the input.536 \n- **Task description localization**: where is the task description blended with the inputs? In the beginning, in the middle, at the end?537 \n- **Implicit situation or contextualization**: how explicit is the query? For instance, *Given this review, would you buy this product?* is an indirect way to ask whether the review is positive.538 """539 540 col1, _, _ = st.beta_columns([18, 1, 6])541 with col1:542 if state.template_name is not None:543 show_text(variety_guideline, with_markdown=True)544 545 #546 # Edit the created or selected template547 #548 col1, _, col2 = st.beta_columns([18, 1, 6])549 with col1:550 if state.template_name is not None:551 template = dataset_templates[state.template_name]552 #553 # If template is selected, displays template editor554 #555 with st.form("edit_template_form"):556 updated_template_name = st.text_input("Name", value=template.name)557 state.reference = st.text_input(558 "Prompt Reference",559 help="Short description of the prompt and/or paper reference for the prompt.",560 value=template.reference,561 )562 563 # Metadata564 state.metadata = template.metadata565 state.metadata.original_task = st.checkbox(566 "Original Task?",567 value=template.metadata.original_task,568 help="Prompt asks model to perform the original task designed for this dataset.",569 )570 state.metadata.choices_in_prompt = st.checkbox(571 "Choices in Template?",572 value=template.metadata.choices_in_prompt,573 help="Prompt explicitly lists choices in the template for the output.",574 )575 576 state.metadata.metrics = st.multiselect(577 "Metrics",578 sorted(METRICS),579 default=template.metadata.metrics,580 help="Select all metrics that are commonly used (or should "581 "be used if a new task) to evaluate this prompt.",582 )583 584 state.metadata.languages = st.multiselect(585 "Prompt Languages",586 sorted(LANGUAGES.keys()),587 default=template.metadata.languages,588 format_func=format_language,589 help="Select all languages used in this prompt. "590 "This annotation is independent from the language(s) "591 "of the dataset.",592 )593 594 # Answer choices595 if template.get_answer_choices_expr() is not None:596 answer_choices = template.get_answer_choices_expr()597 else:598 answer_choices = ""599 state.answer_choices = st.text_input(600 "Answer Choices",601 value=answer_choices,602 help="A Jinja expression for computing answer choices. "603 "Separate choices with a triple bar (|||).",604 )605 606 # Jinja607 state.jinja = st.text_area("Template", height=40, value=template.jinja)608 609 # Submit form610 if st.form_submit_button("Save"):611 if (612 updated_template_name in dataset_templates.all_template_names613 and updated_template_name != state.template_name614 ):615 st.error(616 f"A prompt with the name {updated_template_name} already exists "617 f"for dataset {state.templates_key}."618 )619 elif updated_template_name == "":620 st.error("Need to provide a prompt name.")621 else:622 # Parses state.answer_choices623 if state.answer_choices == "":624 updated_answer_choices = None625 else:626 updated_answer_choices = state.answer_choices627 628 dataset_templates.update_template(629 state.template_name,630 updated_template_name,631 state.jinja,632 state.reference,633 state.metadata,634 updated_answer_choices,635 )636 # Update the state as well637 state.template_name = updated_template_name638 #639 # Displays template output on current example if a template is selected640 # (in second column)641 #642 with col2:643 if state.template_name is not None:644 st.empty()645 template = dataset_templates[state.template_name]646 prompt = template.apply(example)647 if prompt == [""]:648 st.write("β
β
β
*Blank result*")649 else:650 st.write("Input")651 show_text(prompt[0], width=40)652 if len(prompt) > 1:653 st.write("Target")654 show_text(prompt[1], width=40)655 656 #657 # Must sync state at end658 #659 state.sync()660 661 662if __name__ == "__main__":663 run_app()664 