Team Ai
Apppublic

ceyda/ExplaiNER

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
app.py115 linesDownload Raw Back to src
1"""The App module is the main entry point for the application.2 3    Run `streamlit run app.py` to start the app.4"""5 6import pandas as pd7import streamlit as st8from streamlit_option_menu import option_menu9 10from src.load import load_context11from src.subpages import (12    DebugPage,13    FindDuplicatesPage,14    HomePage,15    LossesPage,16    LossySamplesPage,17    MetricsPage,18    MisclassifiedPage,19    Page,20    ProbingPage,21    RandomSamplesPage,22    RawDataPage,23)24from src.subpages.attention import AttentionPage25from src.subpages.hidden_states import HiddenStatesPage26from src.subpages.inspect import InspectPage27from src.utils import classmap28 29sts = st.sidebar30st.set_page_config(31    layout="wide",32    page_title="Error Analysis",33    page_icon="🏷️",34)35 36 37def _show_menu(pages: list[Page]) -> int:38    with st.sidebar:39        page_names = [p.name for p in pages]40        page_icons = [p.icon for p in pages]41        selected_menu_item = st.session_state.active_page = option_menu(42            menu_title="ExplaiNER",43            options=page_names,44            icons=page_icons,45            menu_icon="layout-wtf",46            default_index=0,47        )48        return page_names.index(selected_menu_item)49    assert False50 51 52def _initialize_session_state(pages: list[Page]):53    if "active_page" not in st.session_state:54        for page in pages:55            st.session_state.update(**page._get_widget_defaults())56    st.session_state.update(st.session_state)57 58 59def _write_color_legend(context):60    def style(x):61        return [f"background-color: {rgb}; opacity: 1;" for rgb in colors]62 63    labels = list(set([lbl.split("-")[1] if "-" in lbl else lbl for lbl in context.labels]))64    colors = [st.session_state.get(f"color_{lbl}", "#000000") for lbl in labels]65 66    color_legend_df = pd.DataFrame(67        [classmap[l] for l in labels], columns=["label"], index=labels68    ).T69    st.sidebar.write(70        color_legend_df.T.style.apply(style, axis=0).set_properties(71            **{"color": "white", "text-align": "center"}72        )73    )74 75 76def main():77    """The main entry point for the application."""78    pages: list[Page] = [79        HomePage(),80        AttentionPage(),81        HiddenStatesPage(),82        ProbingPage(),83        MetricsPage(),84        LossySamplesPage(),85        LossesPage(),86        MisclassifiedPage(),87        RandomSamplesPage(),88        FindDuplicatesPage(),89        InspectPage(),90        RawDataPage(),91        DebugPage(),92    ]93 94    _initialize_session_state(pages)95 96    selected_page_idx = _show_menu(pages)97    selected_page = pages[selected_page_idx]98 99    if isinstance(selected_page, HomePage):100        selected_page.render()101        return102 103    if "model_name" not in st.session_state:104        # this can happen if someone loads another page directly (without going through home)105        st.error("Setup not complete. Please click on 'Home / Setup in left menu bar'")106        return107 108    context = load_context(**st.session_state)109    _write_color_legend(context)110    selected_page.render(context)111 112 113if __name__ == "__main__":114    main()115