mrm8488/PromptSource
4
1import argparse2import textwrap3from multiprocessing import Manager, Pool4 5import pandas as pd6import plotly.express as px7import streamlit as st8from datasets import get_dataset_infos9from pygments import highlight10from pygments.formatters import HtmlFormatter11from pygments.lexers import DjangoLexer12 13from session import _get_state14from templates import Template, TemplateCollection15from utils import (16 get_dataset,17 get_dataset_confs,18 list_datasets,19 removeHyphen,20 renameDatasetColumn,21 render_features,22)23 24 25# add an argument for read-only26# At the moment, streamlit does not handle python script arguments gracefully.27# Thus, for read-only mode, you have to type one of the below two:28# streamlit run promptsource/app.py -- -r29# streamlit run promptsource/app.py -- --read-only30# Check https://github.com/streamlit/streamlit/issues/337 for more information.31parser = argparse.ArgumentParser(description="run app.py with args")32parser.add_argument("-r", "--read-only", action="store_true", help="whether to run it as read-only mode")33 34args = parser.parse_args()35if args.read_only:36 select_options = ["Helicopter view", "Prompted dataset viewer"]37 side_bar_title_prefix = "Promptsource (Read only)"38else:39 select_options = ["Helicopter view", "Prompted dataset viewer", "Sourcing"]40 side_bar_title_prefix = "Promptsource"41 42#43# Helper functions for datasets library44#45get_dataset = st.cache(allow_output_mutation=True)(get_dataset)46get_dataset_confs = st.cache(get_dataset_confs)47 48 49def reset_template_state():50 state.template_name = None51 state.jinja = None52 state.reference = None53 54 55#56# Loads session state57#58state = _get_state()59 60#61# Initial page setup62#63st.set_page_config(page_title="Promptsource", layout="wide")64st.sidebar.markdown(65 "<center><a href='https://github.com/bigscience-workshop/promptsource'>💻Github - Promptsource\n\n</a></center>",66 unsafe_allow_html=True,67)68mode = st.sidebar.selectbox(69 label="Choose a mode",70 options=select_options,71 index=0,72 key="mode_select",73)74st.sidebar.title(f"{side_bar_title_prefix} 🌸 - {mode}")75 76#77# Adds pygments styles to the page.78#79st.markdown(80 "<style>" + HtmlFormatter(style="friendly").get_style_defs(".highlight") + "</style>", unsafe_allow_html=True81)82 83WIDTH = 8084 85 86def show_jinja(t, width=WIDTH):87 wrap = textwrap.fill(t, width=width, replace_whitespace=False)88 out = highlight(wrap, DjangoLexer(), HtmlFormatter())89 st.write(out, unsafe_allow_html=True)90 91 92def show_text(t, width=WIDTH, with_markdown=False):93 wrap = [textwrap.fill(subt, width=width, replace_whitespace=False) for subt in t.split("\n")]94 wrap = "\n".join(wrap)95 if with_markdown:96 st.write(wrap, unsafe_allow_html=True)97 else:98 st.text(wrap)99 100 101#102# Loads template data103#104try:105 template_collection = TemplateCollection()106except FileNotFoundError:107 st.error(108 "Unable to find the prompt folder!\n\n"109 "We expect the folder to be in the working directory. "110 "You might need to restart the app in the root directory of the repo."111 )112 st.stop()113 114 115if mode == "Helicopter view":116 st.title("High level metrics")117 st.write(118 "If you want to contribute, please refer to the instructions in "119 + "[Contributing](https://github.com/bigscience-workshop/promptsource/blob/main/CONTRIBUTING.md)."120 )121 122 #123 # Global metrics124 #125 counts = template_collection.get_templates_count()126 nb_prompted_datasets = len(counts)127 st.write(f"## Number of *prompted datasets*: `{nb_prompted_datasets}`")128 nb_prompts = sum(counts.values())129 st.write(f"## Number of *prompts*: `{nb_prompts}`")130 131 #132 # Metrics per dataset/subset133 #134 # Download dataset infos (multiprocessing download)135 manager = Manager()136 all_infos = manager.dict()137 all_datasets = list(set([t[0] for t in template_collection.keys]))138 139 def get_infos(d_name):140 all_infos[d_name] = get_dataset_infos(d_name)141 142 pool = Pool(processes=len(all_datasets))143 pool.map(get_infos, all_datasets)144 pool.close()145 pool.join()146 147 results = []148 for (dataset_name, subset_name) in template_collection.keys:149 # Collect split sizes (train, validation and test)150 if dataset_name not in all_infos:151 infos = get_dataset_infos(dataset_name)152 all_infos[dataset_name] = infos153 else:154 infos = all_infos[dataset_name]155 if infos:156 if subset_name is None:157 subset_infos = infos[list(infos.keys())[0]]158 else:159 subset_infos = infos[subset_name]160 161 split_sizes = {k: v.num_examples for k, v in subset_infos.splits.items()}162 else:163 # Zaid/coqa_expanded and Zaid/quac_expanded don't have dataset_infos.json164 # so infos is an empty dic, and `infos[list(infos.keys())[0]]` raises an error165 # For simplicity, just filling `split_sizes` with nothing, so the displayed split sizes will be 0.166 split_sizes = {}167 168 # Collect template counts, original task counts and names169 dataset_templates = template_collection.get_dataset(dataset_name, subset_name)170 results.append(171 {172 "Dataset name": dataset_name,173 "Subset name": "∅" if subset_name is None else subset_name,174 "Train size": split_sizes["train"] if "train" in split_sizes else 0,175 "Validation size": split_sizes["validation"] if "validation" in split_sizes else 0,176 "Test size": split_sizes["test"] if "test" in split_sizes else 0,177 "Number of prompts": len(dataset_templates),178 "Number of original task prompts": sum(179 [bool(t.metadata.original_task) for t in dataset_templates.templates.values()]180 ),181 "Prompt names": [t.name for t in dataset_templates.templates.values()],182 }183 )184 results_df = pd.DataFrame(results)185 results_df.sort_values(["Number of prompts"], inplace=True, ascending=False)186 results_df.reset_index(drop=True, inplace=True)187 188 nb_training_instances = results_df["Train size"].sum()189 st.write(f"## Number of *training instances*: `{nb_training_instances}`")190 191 plot_df = results_df[["Dataset name", "Subset name", "Train size", "Number of prompts"]].copy()192 plot_df["Name"] = plot_df["Dataset name"] + " - " + plot_df["Subset name"]193 plot_df.sort_values(["Train size"], inplace=True, ascending=False)194 fig = px.bar(195 plot_df,196 x="Name",197 y="Train size",198 hover_data=["Dataset name", "Subset name", "Number of prompts"],199 log_y=True,200 title="Number of training instances per data(sub)set - y-axis is in logscale",201 )202 fig.update_xaxes(visible=False, showticklabels=False)203 st.plotly_chart(fig, use_container_width=True)204 st.write(205 f"- Top 3 training subsets account for `{100*plot_df[:3]['Train size'].sum()/nb_training_instances:.2f}%` of the training instances."206 )207 biggest_training_subset = plot_df.iloc[0]208 st.write(209 f"- Biggest training subset is *{biggest_training_subset['Name']}* with `{biggest_training_subset['Train size']}` instances"210 )211 smallest_training_subset = plot_df[plot_df["Train size"] > 0].iloc[-1]212 st.write(213 f"- Smallest training subset is *{smallest_training_subset['Name']}* with `{smallest_training_subset['Train size']}` instances"214 )215 216 st.markdown("***")217 st.write("Details per dataset")218 st.table(results_df)219 220else:221 # Combining mode `Prompted dataset viewer` and `Sourcing` since the222 # backbone of the interfaces is the same223 assert mode in ["Prompted dataset viewer", "Sourcing"], ValueError(224 f"`mode` ({mode}) should be in `[Helicopter view, Prompted dataset viewer, Sourcing]`"225 )226 227 #228 # Loads dataset information229 #230 231 dataset_list = list_datasets(232 template_collection,233 state,234 )235 ag_news_index = dataset_list.index("ag_news")236 237 #238 # Select a dataset - starts with ag_news239 #240 dataset_key = st.sidebar.selectbox(241 "Dataset",242 dataset_list,243 key="dataset_select",244 index=ag_news_index,245 help="Select the dataset to work on.",246 )247 248 #249 # If a particular dataset is selected, loads dataset and template information250 #251 if dataset_key is not None:252 253 #254 # Check for subconfigurations (i.e. subsets)255 #256 configs = get_dataset_confs(dataset_key)257 conf_option = None258 if len(configs) > 0:259 conf_option = st.sidebar.selectbox("Subset", configs, index=0, format_func=lambda a: a.name)260 261 dataset = get_dataset(dataset_key, str(conf_option.name) if conf_option else None)262 splits = list(dataset.keys())263 index = 0264 if "train" in splits:265 index = splits.index("train")266 split = st.sidebar.selectbox("Split", splits, key="split_select", index=index)267 dataset = dataset[split]268 dataset = renameDatasetColumn(dataset)269 270 dataset_templates = template_collection.get_dataset(dataset_key, conf_option.name if conf_option else None)271 272 template_list = dataset_templates.all_template_names273 num_templates = len(template_list)274 st.sidebar.write(275 "No of prompts created for "276 + f"`{dataset_key + (('/' + conf_option.name) if conf_option else '')}`"277 + f": **{str(num_templates)}**"278 )279 280 if mode == "Prompted dataset viewer":281 if num_templates > 0:282 template_name = st.sidebar.selectbox(283 "Prompt name",284 template_list,285 key="template_select",286 index=0,287 help="Select the prompt to visualize.",288 )289 290 step = 50291 example_index = st.sidebar.number_input(292 f"Select the example index (Size = {len(dataset)})",293 min_value=0,294 max_value=len(dataset) - step,295 value=0,296 step=step,297 key="example_index_number_input",298 help="Offset = 50.",299 )300 else: # mode = Sourcing301 st.sidebar.subheader("Select Example")302 example_index = st.sidebar.slider("Select the example index", 0, len(dataset) - 1)303 304 example = dataset[example_index]305 example = removeHyphen(example)306 307 st.sidebar.write(example)308 309 st.sidebar.subheader("Dataset Schema")310 rendered_features = render_features(dataset.features)311 st.sidebar.write(rendered_features)312 313 #314 # Display dataset information315 #316 st.header("Dataset: " + dataset_key + " " + (("/ " + conf_option.name) if conf_option else ""))317 318 st.markdown(319 "*Homepage*: "320 + dataset.info.homepage321 + "\n\n*Dataset*: https://github.com/huggingface/datasets/blob/master/datasets/%s/%s.py"322 % (dataset_key, dataset_key)323 )324 325 md = """326 %s327 """ % (328 dataset.info.description.replace("\\", "") if dataset_key else ""329 )330 st.markdown(md)331 332 #333 # Body of the app: display prompted examples in mode `Prompted dataset viewer`334 # or text boxes to create new prompts in mode `Sourcing`335 #336 if mode == "Prompted dataset viewer":337 #338 # Display template information339 #340 if num_templates > 0:341 template = dataset_templates[template_name]342 st.subheader("Prompt")343 st.markdown("##### Name")344 st.text(template.name)345 st.markdown("##### Reference")346 st.text(template.reference)347 st.markdown("##### Original Task? ")348 st.text(template.metadata.original_task)349 st.markdown("##### Choices in template? ")350 st.text(template.metadata.choices_in_prompt)351 st.markdown("##### Metrics")352 st.text(", ".join(template.metadata.metrics) if template.metadata.metrics else None)353 st.markdown("##### Answer Choices")354 if template.get_answer_choices_expr() is not None:355 show_jinja(template.get_answer_choices_expr())356 else:357 st.text(None)358 st.markdown("##### Jinja template")359 splitted_template = template.jinja.split("|||")360 st.markdown("###### Input template")361 show_jinja(splitted_template[0].strip())362 if len(splitted_template) > 1:363 st.markdown("###### Target template")364 show_jinja(splitted_template[1].strip())365 st.markdown("***")366 367 #368 # Display a couple (steps) examples369 #370 for ex_idx in range(example_index, example_index + step):371 if ex_idx >= len(dataset):372 continue373 example = dataset[ex_idx]374 example = removeHyphen(example)375 col1, _, col2 = st.beta_columns([12, 1, 12])376 with col1:377 st.write(example)378 if num_templates > 0:379 with col2:380 prompt = template.apply(example, highlight_variables=False)381 if prompt == [""]:382 st.write("∅∅∅ *Blank result*")383 else:384 st.write("Input")385 show_text(prompt[0])386 if len(prompt) > 1:387 st.write("Target")388 show_text(prompt[1])389 st.markdown("***")390 else: # mode = Sourcing391 st.markdown("## Prompt Creator")392 393 #394 # Create a new template or select an existing one395 #396 col1a, col1b, _, col2 = st.beta_columns([9, 9, 1, 6])397 398 # current_templates_key and state.templates_key are keys for the templates object399 current_templates_key = (dataset_key, conf_option.name if conf_option else None)400 401 # Resets state if there has been a change in templates_key402 if state.templates_key != current_templates_key:403 state.templates_key = current_templates_key404 reset_template_state()405 406 with col1a, st.form("new_template_form"):407 new_template_name = st.text_input(408 "Create a New Prompt",409 key="new_template",410 value="",411 help="Enter name and hit enter to create a new prompt.",412 )413 new_template_submitted = st.form_submit_button("Create")414 if new_template_submitted:415 if new_template_name in dataset_templates.all_template_names:416 st.error(417 f"A prompt with the name {new_template_name} already exists "418 f"for dataset {state.templates_key}."419 )420 elif new_template_name == "":421 st.error("Need to provide a prompt name.")422 else:423 template = Template(new_template_name, "", "")424 dataset_templates.add_template(template)425 reset_template_state()426 state.template_name = new_template_name427 else:428 state.new_template_name = None429 430 with col1b, st.beta_expander("or Select Prompt", expanded=True):431 dataset_templates = template_collection.get_dataset(*state.templates_key)432 template_list = dataset_templates.all_template_names433 if state.template_name:434 index = template_list.index(state.template_name)435 else:436 index = 0437 state.template_name = st.selectbox(438 "", template_list, key="template_select", index=index, help="Select the prompt to work on."439 )440 441 if st.button("Delete Prompt", key="delete_prompt"):442 dataset_templates.remove_template(state.template_name)443 reset_template_state()444 445 variety_guideline = """446 :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.447 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.448 \r**To get various prompts, you can try moving the cursor along theses axes**:449 \n- **Interrogative vs affirmative form**: Ask a question about an attribute of the inputs or tell the model to decide something about the input.450 \n- **Task description localization**: where is the task description blended with the inputs? In the beginning, in the middle, at the end?451 \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.452 """453 454 col1, _, _ = st.beta_columns([18, 1, 6])455 with col1:456 if state.template_name is not None:457 show_text(variety_guideline, with_markdown=True)458 459 #460 # Edit the created or selected template461 #462 col1, _, col2 = st.beta_columns([18, 1, 6])463 with col1:464 if state.template_name is not None:465 template = dataset_templates[state.template_name]466 #467 # If template is selected, displays template editor468 #469 with st.form("edit_template_form"):470 updated_template_name = st.text_input("Name", value=template.name)471 state.reference = st.text_input(472 "Prompt Reference",473 help="Short description of the prompt and/or paper reference for the prompt.",474 value=template.reference,475 )476 477 # Metadata478 state.metadata = template.metadata479 state.metadata.original_task = st.checkbox(480 "Original Task?",481 value=template.metadata.original_task,482 help="Prompt asks model to perform the original task designed for this dataset.",483 )484 state.metadata.choices_in_prompt = st.checkbox(485 "Choices in Template?",486 value=template.metadata.choices_in_prompt,487 help="Prompt explicitly lists choices in the template for the output.",488 )489 490 # Metrics from here:491 # https://github.com/google-research/text-to-text-transfer-transformer/blob/4b580f23968c2139be7fb1cd53b22c7a7f686cdf/t5/evaluation/metrics.py492 metrics_choices = [493 "BLEU",494 "ROUGE",495 "Squad",496 "Trivia QA",497 "Accuracy",498 "Pearson Correlation",499 "Spearman Correlation",500 "MultiRC",501 "AUC",502 "COQA F1",503 "Edit Distance",504 ]505 # Add mean reciprocal rank506 metrics_choices.append("Mean Reciprocal Rank")507 # Add generic other508 metrics_choices.append("Other")509 # Sort alphabetically510 metrics_choices = sorted(metrics_choices)511 state.metadata.metrics = st.multiselect(512 "Metrics",513 metrics_choices,514 default=template.metadata.metrics,515 help="Select all metrics that are commonly used (or should "516 "be used if a new task) to evaluate this prompt.",517 )518 519 # Answer choices520 if template.get_answer_choices_expr() is not None:521 answer_choices = template.get_answer_choices_expr()522 else:523 answer_choices = ""524 state.answer_choices = st.text_input(525 "Answer Choices",526 value=answer_choices,527 help="A Jinja expression for computing answer choices. "528 "Separate choices with a triple bar (|||).",529 )530 531 # Jinja532 state.jinja = st.text_area("Template", height=40, value=template.jinja)533 534 # Submit form535 if st.form_submit_button("Save"):536 if (537 updated_template_name in dataset_templates.all_template_names538 and updated_template_name != state.template_name539 ):540 st.error(541 f"A prompt with the name {updated_template_name} already exists "542 f"for dataset {state.templates_key}."543 )544 elif updated_template_name == "":545 st.error("Need to provide a prompt name.")546 else:547 # Parses state.answer_choices548 if state.answer_choices == "":549 updated_answer_choices = None550 else:551 updated_answer_choices = state.answer_choices552 553 dataset_templates.update_template(554 state.template_name,555 updated_template_name,556 state.jinja,557 state.reference,558 state.metadata,559 updated_answer_choices,560 )561 # Update the state as well562 state.template_name = updated_template_name563 #564 # Displays template output on current example if a template is selected565 # (in second column)566 #567 with col2:568 if state.template_name is not None:569 st.empty()570 template = dataset_templates[state.template_name]571 prompt = template.apply(example)572 if prompt == [""]:573 st.write("∅∅∅ *Blank result*")574 else:575 st.write("Input")576 show_text(prompt[0], width=40)577 if len(prompt) > 1:578 st.write("Target")579 show_text(prompt[1], width=40)580 581 582#583# Must sync state at end584#585state.sync()586 