Team Ai
Apppublic

mrm8488/PromptSource

sourceHugging Faceupdated 5y agoView on Hugging Face
4likes
app.py586 linesDownload Raw Back to root
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