Team Ai
Apppublic

autoevaluate/error-analysis

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
app.py283 linesDownload Raw Back to root
1## LIBRARIES ###2## Data3import numpy as np4import pandas as pd5import torch6import json7from tqdm import tqdm8from math import floor9from datasets import load_dataset10from collections import defaultdict11from transformers import AutoTokenizer12pd.options.display.float_format = '${:,.2f}'.format13 14# Analysis15# from gensim.models.doc2vec import Doc2Vec16# from sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score17import nltk18from nltk.cluster import KMeansClusterer19import scipy.spatial.distance as sdist20from scipy.spatial import distance_matrix21# nltk.download('punkt') #make sure that punkt is downloaded22 23# App & Visualization24import streamlit as st25import altair as alt26import plotly.graph_objects as go27from streamlit_vega_lite import altair_component28 29 30 31# utils32from random import sample33from error_analysis import utils as ut34 35 36def down_samp(embedding):37    """Down sample a data frame for altiar visualization """38    # total number of positive and negative sentiments in the class39    #embedding = embedding.groupby('slice').apply(lambda x: x.sample(frac=0.3))40    total_size = embedding.groupby(['slice','label'], as_index=False).count()41 42    user_data = 043    # if 'Your Sentences' in str(total_size['slice']):44    #     tmp = embedding.groupby(['slice'], as_index=False).count()45    #     val = int(tmp[tmp['slice'] == "Your Sentences"]['source'])46    #     user_data = val47 48    max_sample = total_size.groupby('slice').max()['content']49 50    # # down sample to meeting altair's max values51    # # but keep the proportional representation of groups52    down_samp = 1/(sum(max_sample.astype(float))/(1000-user_data))53 54    max_samp = max_sample.apply(lambda x: floor(x*down_samp)).astype(int).to_dict()55    max_samp['Your Sentences'] = user_data56 57    # # sample down for each group in the data frame58    embedding = embedding.groupby('slice').apply(lambda x: x.sample(n=max_samp.get(x.name))).reset_index(drop=True)59 60    # # order the embedding61    return(embedding)62 63 64def data_comparison(df):65    selection = alt.selection_multi(fields=['cluster:N','label:O'])66    color = alt.condition(alt.datum.slice == 'high-loss', alt.Color('cluster:N', scale = alt.Scale(domain=df.cluster.unique().tolist())), alt.value("lightgray"))67    opacity = alt.condition(selection, alt.value(0.7), alt.value(0.25))68 69    # basic chart70    scatter = alt.Chart(df).mark_point(size=100, filled=True).encode(71        x=alt.X('x:Q', axis=None),72        y=alt.Y('y:Q', axis=None),73        color=color,74        shape=alt.Shape('label:O', scale=alt.Scale(range=['circle', 'diamond'])),75        tooltip=['cluster:N','slice:N','content:N','label:O','pred:O'],76        opacity=opacity77    ).properties(78        width=1000,79        height=80080    ).interactive()81 82    legend = alt.Chart(df).mark_point(size=100, filled=True).encode(83        x=alt.X("label:O"),84        y=alt.Y('cluster:N', axis=alt.Axis(orient='right'), title=""),85        shape=alt.Shape('label:O', scale=alt.Scale(86        range=['circle', 'diamond']), legend=None),87        color=color,88    ).add_selection(89        selection90    )91    layered = scatter | legend92    layered = layered.configure_axis(93        grid=False94    ).configure_view(95        strokeOpacity=096    )97    return layered98 99def quant_panel(embedding_df):100    """ Quantitative Panel Layout"""101    all_metrics = {}102    st.warning("**Error slice visualization**")103    with st.expander("How to read this chart:"):104        st.markdown("* Each **point** is an input example.")105        st.markdown("* Gray points have low-loss and the colored have high-loss. High-loss instances are clustered using **kmeans** and each color represents a cluster.")106        st.markdown("* The **shape** of each point reflects the label category --  positive (diamond) or negative sentiment (circle).")107    st.altair_chart(data_comparison(down_samp(embedding_df)), use_container_width=True)108 109 110def frequent_tokens(data, tokenizer, loss_quantile=0.95, top_k=200, smoothing=0.005):111    unique_tokens = []112    tokens = []113    for row in tqdm(data['content']):114        tokenized = tokenizer(row,padding=True, return_tensors='pt')115        tokens.append(tokenized['input_ids'].flatten())116        unique_tokens.append(torch.unique(tokenized['input_ids']))117    losses = data['loss'].astype(float)118    high_loss = losses.quantile(loss_quantile)119    loss_weights = (losses > high_loss)120    loss_weights = loss_weights / loss_weights.sum()121    token_frequencies = defaultdict(float)122    token_frequencies_error = defaultdict(float)123 124    weights_uniform = np.full_like(loss_weights, 1 / len(loss_weights))125 126    num_examples = len(data)127    for i in tqdm(range(num_examples)):128        for token in unique_tokens[i]:129            token_frequencies[token.item()] += weights_uniform[i]130            token_frequencies_error[token.item()] += loss_weights[i]131 132    token_lrs = {k: (smoothing+token_frequencies_error[k]) / (smoothing+token_frequencies[k]) for k in token_frequencies}133    tokens_sorted = list(map(lambda x: x[0], sorted(token_lrs.items(), key=lambda x: x[1])[::-1]))134 135    top_tokens = []136    for i, (token) in enumerate(tokens_sorted[:top_k]):137        top_tokens.append(['%10s' % (tokenizer.decode(token)), '%.4f' % (token_frequencies[token]), '%.4f' % (138            token_frequencies_error[token]), '%4.2f' % (token_lrs[token])])139    return pd.DataFrame(top_tokens, columns=['Token', 'Freq', 'Freq error slice', 'lrs'])140 141 142@st.cache(ttl=600)143def get_data(inference, emb):144    preds = inference.outputs.numpy()145    losses = inference.losses.numpy()146    embeddings = pd.DataFrame(emb, columns=['x', 'y'])147    num_examples = len(losses)148    # dataset_labels = [dataset[i]['label'] for i in range(num_examples)]149    return pd.concat([pd.DataFrame(np.transpose(np.vstack([dataset[:num_examples]['content'], 150                    dataset[:num_examples]['label'], preds, losses])), columns=['content', 'label', 'pred', 'loss']), embeddings], axis=1)151 152def clustering(data,num_clusters):153    X = np.array(data['embedding'].tolist())154    kclusterer = KMeansClusterer(155        num_clusters, distance=nltk.cluster.util.cosine_distance,156        repeats=25,avoid_empty_clusters=True)157    assigned_clusters = kclusterer.cluster(X, assign_clusters=True)158    data['cluster'] = pd.Series(assigned_clusters, index=data.index).astype('int')159    data['centroid'] = data['cluster'].apply(lambda x: kclusterer.means()[x])160    return data, assigned_clusters161 162def kmeans(df, num_clusters=3):163    data_hl = df.loc[df['slice'] == 'high-loss']164    data_kmeans,clusters = clustering(data_hl,num_clusters)165    merged = pd.merge(df, data_kmeans, left_index=True, right_index=True, how='outer', suffixes=('', '_y'))166    merged.drop(merged.filter(regex='_y$').columns.tolist(),axis=1,inplace=True)167    merged['cluster'] = merged['cluster'].fillna(num_clusters).astype('int')168    return merged169 170def distance_from_centroid(row):171    return sdist.norm(row['embedding'] - row['centroid'].tolist())172 173@st.cache(ttl=600)174def topic_distribution(weights, smoothing=0.01):175    topic_frequencies = defaultdict(float)176    topic_frequencies_spotlight = defaultdict(float)177    weights_uniform = np.full_like(weights, 1 / len(weights))178    num_examples = len(weights)179    for i in range(num_examples):180        example = dataset[i]181        category = example['title']182        topic_frequencies[category] += weights_uniform[i]183        topic_frequencies_spotlight[category] += weights[i]184 185    topic_ratios = {c: (smoothing + topic_frequencies_spotlight[c]) / (186        smoothing + topic_frequencies[c]) for c in topic_frequencies}187 188    categories_sorted = map(lambda x: x[0], sorted(189        topic_ratios.items(), key=lambda x: x[1], reverse=True))190 191    topic_distr = []192    for category in categories_sorted:193        topic_distr.append(['%.3f' % topic_frequencies[category], '%.3f' %194                           topic_frequencies_spotlight[category], '%.2f' % topic_ratios[category], '%s' % category])195 196    return pd.DataFrame(topic_distr, columns=['Overall frequency', 'Error frequency', 'Ratio', 'Category'])197    # for category in categories_sorted:198    #    return(topic_frequencies[category], topic_frequencies_spotlight[category], topic_ratios[category], category)199 200def populate_session(dataset,model):201    data_df = read_file_to_df('./assets/data/'+dataset+ '_'+ model+'.parquet')202    if model == 'albert-base-v2-yelp-polarity':203        tokenizer = AutoTokenizer.from_pretrained('textattack/'+model)204    else:205        tokenizer = AutoTokenizer.from_pretrained(model)206    if "user_data" not in st.session_state:207        st.session_state["user_data"] = data_df208    if "selected_slice" not in st.session_state:209        st.session_state["selected_slice"] = None210 211@st.cache(allow_output_mutation=True)212def read_file_to_df(file):213   return pd.read_parquet(file)214 215if __name__ == "__main__":216    ### STREAMLIT APP CONGFIG ###217    st.set_page_config(layout="wide", page_title="Interactive Error Analysis")218 219    ut.init_style()220 221    lcol, rcol = st.columns([2, 2])222    # ******* loading the mode and the data223    #st.sidebar.mardown("<h4>Interactive Error Analysis</h4>", unsafe_allow_html=True)224 225    dataset = st.sidebar.selectbox(226        "Dataset",227        ["amazon_polarity", "yelp_polarity"],228        index = 1229    )230 231    model = st.sidebar.selectbox(232        "Model",233        ["distilbert-base-uncased-finetuned-sst-2-english",234            "albert-base-v2-yelp-polarity"],235    )236 237    ### LOAD DATA AND SESSION VARIABLES ###238    ##uncomment the next next line to run dynamically and not from file239    #populate_session(dataset, model)240    data_df = read_file_to_df('./assets/data/'+dataset+ '_'+ model+'.parquet')241    loss_quantile = st.sidebar.slider(242        "Loss Quantile", min_value=0.5, max_value=1.0,step=0.01,value=0.95243    )244    data_df['loss'] = data_df['loss'].astype(float)245    losses = data_df['loss']246    high_loss = losses.quantile(loss_quantile)247    data_df['slice'] = 'high-loss'248    data_df['slice'] = data_df['slice'].where(data_df['loss'] > high_loss, 'low-loss') 249 250    with rcol:251        with st.spinner(text='loading...'):252            st.markdown('<h3>Word Distribution in Error Slice</h3>', unsafe_allow_html=True)253            #uncomment the next two lines to run dynamically and not from file254            #commontokens = frequent_tokens(data_df, tokenizer, loss_quantile=loss_quantile)255            commontokens = read_file_to_df('./assets/data/'+dataset+ '_'+ model+'_commontokens.parquet')256            with st.expander("How to read the table:"):257                st.markdown("* The table displays the most frequent tokens in error slices, relative to their frequencies in the val set.")258            st.write(commontokens)259 260    run_kmeans = st.sidebar.radio("Cluster error slice?", ('True', 'False'), index=0)261 262    num_clusters = st.sidebar.slider("# clusters", min_value=1, max_value=20, step=1, value=3)263 264    if run_kmeans == 'True':265        with st.spinner(text='running kmeans...'):266            merged = kmeans(data_df,num_clusters=num_clusters)267    with lcol:268        st.markdown('<h3>Error Slices</h3>',unsafe_allow_html=True)269        with st.expander("How to read the table:"):270            st.markdown("* *Error slice* refers to the subset of evaluation dataset the model performs poorly on.")271            st.markdown("* The table displays model error slices on the evaluation dataset, sorted by loss.")272            st.markdown("* Each row is an input example that includes the label, model pred, loss, and error cluster.")273        with st.spinner(text='loading error slice...'):274            dataframe=read_file_to_df('./assets/data/'+dataset+ '_'+ model+'_error-slices.parquet')275        #uncomment the next next line to run dynamically and not from file276        # dataframe = merged[['content', 'label', 'pred', 'loss', 'cluster']].sort_values(277        #     by=['loss'], ascending=False)278        # table_html = dataframe.to_html(279        #     columns=['content', 'label', 'pred', 'loss', 'cluster'], max_rows=50)280        # table_html = table_html.replace("<th>", '<th align="left">')  # left-align the headers281            st.write(dataframe,width=900, height=300)282    with st.spinner(text='loading visualization...'):283        quant_panel(merged)