Team Ai
Apppublic

bigscience/promptsource

sourceHugging Faceupdated 3y agoView on Hugging Face
105likes
app.py664 linesDownload Raw Back to promptsource
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