Team Ai
Apppublic

GAIR/Preference-Dissection-Visualization

sourceHugging Facemitupdated 3y agoView on Hugging Face
6likes
app.py941 linesDownload Raw Back to root
1import streamlit as st2import os3from utils import read_all, json_to_markdown_bold_keys, custom_md_with_color4from scipy.stats import pearsonr, spearmanr5import seaborn as sns6import pandas as pd7import json8 9import jax10import jax.numpy as jnp11import numpy as np12import numpyro13import numpyro.distributions as dist14from numpyro.infer import MCMC, NUTS15from matplotlib import pyplot as plt16import shap17from functools import partial18 19 20import base6421 22numpyro.set_host_device_count(4)23 24feature_name_to_id = {25    "harmlessness": 0,26    "grammar, spelling, punctuation, and code-switching": 1,27    "friendly": 2,28    "polite": 3,29    "interactive": 4,30    "authoritative tone": 5,31    "funny and humorous": 6,32    "metaphors, personification, similes, hyperboles, irony, parallelism": 7,33    "complex word usage and sentence structure": 8,34    "use of direct and explicit supporting materials": 9,35    "well formatted": 10,36    "admit limitations or mistakes": 11,37    "persuade user": 12,38    "step by step solution": 13,39    "use of informal expressions": 14,40    "non-repetitive": 15,41    "clear and understandable": 16,42    "relevance without considering inaccuracy": 17,43    "innovative and novel": 18,44    "information richness without considering inaccuracy": 19,45    "no minor errors": 20,46    "no moderate errors": 21,47    "no severe errors": 22,48    "clarify user intent": 23,49    "showing empathetic": 24,50    "satisfying explicit constraints": 25,51    "supporting explicit subjective stances": 26,52    "correcting explicit mistakes or biases": 27,53    "length": 28,54}55 56feature_name_to_id_short = {57    "harmless": 0,58    "grammarly correct": 1,59    "friendly": 2,60    "polite": 3,61    "interactive": 4,62    "authoritative": 5,63    "funny": 6,64    "use rhetorical devices": 7,65    "complex word & sentence": 8,66    "use supporting materials": 9,67    "well formatted": 10,68    "admit limits": 11,69    "persuasive": 12,70    "step-by-step": 13,71    "use informal expressions": 14,72    "non-repetitive": 15,73    "clear": 16,74    "relevant": 17,75    "novel": 18,76    "contain rich info": 19,77    "no minor errors": 20,78    "no moderate errors": 21,79    "no severe errors": 22,80    "clarify intent": 23,81    "show empathetic": 24,82    "satisfy constraints": 25,83    "support stances": 26,84    "correct mistakes": 27,85    "lengthy": 28,86}87 88small_mapping_for_query_specific_cases = {89    "w_constraints": "Contain Explicit Constraints",90    "w_stances": "Show Explicit Subjective Stances",91    "w_mistakes": "Contain Mistakes or Bias",92    "intent_unclear": "Unclear User Intent",93    "express_feeling": "Express Feelings of Emotions",94}95 96pre_set_full_model_order = [97    "yi-6b",98    "yi-6b-chat",99    "llama-2-7b",100    "llama-2-7b-chat",101    "vicuna-7b-v1.5",102    "tulu-2-dpo-7b",103    "mistral-7b",104    "mistral-7b-instruct-v0.1",105    "mistral-7b-instruct-v0.2",106    "zephyr-7b-alpha",107    "zephyr-7b-beta",108    "qwen-7b",109    "qwen-7b-chat",110    "llama-2-13b",111    "llama-2-13b-chat",112    "wizardLM-13b-v1.2",113    "vicuna-13b-v1.5",114    "tulu-2-dpo-13b",115    "qwen-14b",116    "qwen-14b-chat",117    "yi-34b",118    "yi-34b-chat",119    "mistral-8x7b",120    "mistral-8x7b-instruct-v0.1",121    "llama-2-70b",122    "llama-2-70b-chat",123    "wizardLM-70b-v1.0",124    "tulu-2-dpo-70b",125    "qwen-72b",126    "qwen-72b-chat",127    "gpt-3.5-turbo-1106",128    "gpt-4-1106-preview",129    "human",130]131 132feature_id_to_name_short = {v: k for k, v in feature_name_to_id_short.items()}133 134feature_names_short = list(feature_name_to_id_short.keys())135 136all_models_fitted_params = {}137 138def formal_group_name(part):139    if part[0].isupper():140        part = f"[Scenario] {part}"141    else:142        part = f"[Query-Specific Cases] {small_mapping_for_query_specific_cases[part]}"143    return part144 145for fn in os.listdir(f"./data/fitted_paras_comparison"):146    part = fn[len("model_"): fn.find("_fitted_paras")]147    part = formal_group_name(part)148    if part not in all_models_fitted_params:149        all_models_fitted_params[part] = {}150    dd = read_all(f"./data/fitted_paras_comparison/{fn}")151    for it in dd:152        all_models_fitted_params[part][it["model_name"]] = it["parameters"]153 154modelwise_fitted_paras = {}155for group in all_models_fitted_params:156    for model in all_models_fitted_params[group]:157        if model not in modelwise_fitted_paras:158            modelwise_fitted_paras[model] = {}159        modelwise_fitted_paras[model][group] = all_models_fitted_params[group][model]160 161 162def show_one_model_prob(weights, feature_names=None):163    plt.figure(figsize=(20, 7))164 165    plt.rcParams["font.family"] = "Times New Roman"166    plt.rcParams["font.size"] = 20167 168    all_probabilities = []169 170    weights = np.asarray(weights)171    posterior_means = weights172    X_test = np.eye(weights.shape[0])173 174    logits = X_test @ posterior_means175    probabilities = 100 / (1 + np.exp(-logits))176    all_probabilities.extend(probabilities)177 178    plt.scatter(179        range(0, weights.shape[0]),180        probabilities,181        label='apple',182        s=380,183        alpha=0.65,184    )185 186    min_prob = min(all_probabilities)187    max_prob = max(all_probabilities)188    plt.ylim([min_prob - 3, max_prob + 3])189 190    # plt.xlabel('Feature Names')191    plt.ylabel("Probability of Preferred (%)")192    # plt.legend(loc="upper left", bbox_to_anchor=(1, 1))193 194    if feature_names is not None:195        plt.xticks(range(0, len(feature_names)), feature_names, rotation=45, ha="right")196    else:197        plt.xticks(range(0, weights.shape[0]), ha="center")198 199    plt.grid(True)200    plt.axhline(y=50, color="red", linestyle="--")201 202    plt.subplots_adjust(bottom=0.3, right=0.85)203    plt.tight_layout()204    st.pyplot(plt)205    plt.clf()206 207 208def show_all_models_prob(models, selected_models, feature_names=None):209    plt.figure(figsize=(17, 7))210 211    plt.rcParams["font.family"] = "Times New Roman"212    plt.rcParams["font.size"] = 20213 214    all_probabilities = []215    for model_name in selected_models:216        weights = np.asarray(models[model_name])217        posterior_means = weights218        X_test = np.eye(weights.shape[0])219 220        logits = X_test @ posterior_means221        probabilities = 100 / (1 + np.exp(-logits))222        all_probabilities.extend(probabilities)223 224        plt.scatter(225            range(0, weights.shape[0]),226            probabilities,227            label=model_name,228            s=380,229            alpha=0.65,230        )231 232    min_prob = min(all_probabilities)233    max_prob = max(all_probabilities)234    plt.ylim([min_prob - 3, max_prob + 3])235 236    # plt.xlabel('Feature Names')237    plt.ylabel("Probability of Preferred (%)")238    plt.legend(loc="upper left", bbox_to_anchor=(1, 1))239 240    if feature_names is not None:241        plt.xticks(range(0, len(feature_names)), feature_names, rotation=45, ha="right")242    else:243        plt.xticks(range(0, weights.shape[0]), ha="center")244 245    plt.grid(True)246    plt.axhline(y=50, color="red", linestyle="--")247 248    plt.subplots_adjust(bottom=0.3, right=0.85)249    plt.tight_layout()250    st.pyplot(plt)251    plt.clf()252 253 254def process_query_info(x):255    s = []256    if x["clear intent"] != "Yes":257        s.append("[Query-Specific Cases] Unclear User Intent")258    if x["explicitly express feelings"] == "Yes":259        s.append("[Query-Specific Cases] Express Feelings of Emotions")260    if len(x["explicit constraints"]) > 0:261        s.append("[Query-Specific Cases] Contain Explicit Constraints")262    if len(x["explicit subjective stances"]) > 0:263        s.append("[Query-Specific Cases] Show Explicit Subjective Stances")264    if len(x["explicit mistakes or biases"]) > 0:265        s.append("[Query-Specific Cases] Contain Mistakes or Bias")266    return s267 268 269def get_feature(item, remove_length=False, way="comparison"):270    # way be "comparison" or "diff" or "norm_diff"271    feature = [0] * len(feature_name_to_id)272    comparison = item["comparison"]273    for k, v in comparison.items():274        if k == "accuracy":275            for xx in ["Severe", "Moderate", "Minor"]:276                feature[feature_name_to_id[f"no {xx.lower()} errors"]] = v[way][xx]277        elif k == "repetitive":278            feature[feature_name_to_id["non-repetitive"]] = -v[way]279        else:280            feature[feature_name_to_id[k]] = v[way]281    if remove_length:282        feature = feature[:-1]283    return feature284 285 286class BayesianLogisticRegression:287    def __init__(self, alpha):288        self.alpha = alpha289 290    def predict(self, X):291        probs = self.return_prob(X)292        predictions = np.round(probs)293        return predictions294 295    def return_prob(self, X):296        logits = np.dot(X, self.alpha)297        # return probabilities298        return np.exp(logits) / (1 + np.exp(logits))299 300 301def bayesian_logistic_regression(X, y, scale=0.01):302    # Priors for the regression coefficients303    alpha = numpyro.sample('alpha', dist.Laplace(loc=jnp.zeros(X.shape[1]), scale=scale))304 305    # Calculate the linear predictor (the logits) using JAX NumPy306    logits = jnp.dot(X, alpha)307 308    # Likelihood of the observations given the logistic model309    with numpyro.plate('data', X.shape[0]):310        numpyro.sample('obs', dist.Bernoulli(logits=logits), obs=y)311 312 313def fit_bayes_logistic_regression(X, y, scale=0.1, ):314    # repeat X and y on the first axis to get more samples315 316    bxx = partial(bayesian_logistic_regression, scale=scale)317 318    kernel = NUTS(bxx)319    mcmc = MCMC(kernel, num_warmup=500, num_samples=2000, num_chains=4, progress_bar=False)320    mcmc.run(jax.random.PRNGKey(0), X, y)321 322    # Get the posterior samples323    posterior_samples = mcmc.get_samples()324 325    # Compute the mean of the posterior for each alpha_i326    alpha_mean = np.mean(posterior_samples['alpha'], axis=0).tolist()327 328    return BayesianLogisticRegression(alpha_mean), alpha_mean329 330 331def get_similarity(dict1, dict2, type="pearson", select_part="Overall"):332    assert dict1.keys() == dict2.keys(), "Dicts must have the same keys"333    if select_part == "Overall":334        all_sim = 0.0335        count = 0.0336        for key in dict1.keys():337            if key.startswith("[Query-Specific Cases]"): continue338            sim = get_similarity_local(dict1[key], dict2[key], type)339            all_sim += sim340            count += 1341        return all_sim / count342    else:343        return get_similarity_local(dict1[select_part], dict2[select_part], type)344 345 346def get_similarity_local(list1, list2, type="pearson"):347    """348    Calculate the similarity between two lists of numbers based on the specified type.349 350    :param list1: a dict, each field is a list of floats351    :param list2: a dict, each field is a list of floats352    :param type: which kind of 'similarity' is calculated353    :return: the calculated similarity354    """355    assert len(list1) == len(list2), "Lists must be of the same length"356 357    if type == "pearson":358        # Pearson correlation359        similarity, _ = pearsonr(list1, list2)360    elif type == "spearman":361        # Spearman correlation362        similarity, _ = spearmanr(list1, list2)363    elif type == "normed_l1":364        # Normalized negative L1 norm (Manhattan distance)365        similarity = -np.sum(np.abs(np.array(list1) - np.array(list2))) / len(list1)366    elif type == "normed_l2":367        # Normalized negative L2 norm (Euclidean distance)368        similarity = -np.sqrt(np.sum((np.array(list1) - np.array(list2)) ** 2)) / len(369            list1370        )371    else:372        raise NotImplementedError("The specified similarity type is not implemented")373 374    return similarity375 376 377@st.cache_resource378def calculate_similarity_matrix(379        modelwise_fitted_paras, selected_models, similarity_type, selected_part380):381    # Initialize a matrix to store similarities382    if similarity_type in ["spearman", "pearson"]:383        similarity_matrix = np.ones((len(selected_models), len(selected_models)))384    else:385        similarity_matrix = np.zeros((len(selected_models), len(selected_models)))386 387    # Calculate similarities388    for i, model1 in enumerate(selected_models):389        for j, model2 in enumerate(selected_models):390            if i < j:  # Calculate only for upper triangular matrix391                sim = get_similarity(392                    modelwise_fitted_paras[model1],393                    modelwise_fitted_paras[model2],394                    similarity_type,395                    selected_part,396                )397                similarity_matrix[i, j] = sim398                similarity_matrix[j, i] = sim  # Symmetric matrix399    return similarity_matrix400 401 402def format_matrix(matrix):403    formatted_matrix = np.array(matrix, dtype=str)404    for i in range(matrix.shape[0]):405        for j in range(matrix.shape[1]):406            formatted_matrix[i, j] = f"{matrix[i, j]:.2f}".lstrip("0")407    return formatted_matrix408 409 410def become_formal(name):411    name = (412        name.replace("6b", "6B")413        .replace("7b", "7B")414        .replace("13b", "13B")415        .replace("14b", "14B")416        .replace("34b", "34B")417        .replace("70b", "70B")418        .replace("72b", "72B")419    )420    name = (421        name.replace("llama", "LLaMA")422        .replace("yi", "Yi")423        .replace("mistral", "Mistral")424        .replace("qwen", "Qwen")425        .replace("tulu", "Tulu")426        .replace("vicuna", "Vicuna")427        .replace("wizardLM", "WizardLM")428        .replace("zephyr", "Zephyr")429    )430    name = name.replace("chat", "Chat")431    name = name.replace("gpt-3.5-turbo-1106", "GPT-3.5-Turbo").replace(432        "gpt-4-1106-preview", "GPT-4-Turbo"433    )434    name = (435        name.replace("instruct", "Inst").replace("dpo", "DPO").replace("human", "Human")436    )437    return name438 439 440def display_markdown_with_scroll(text, height=200):441    """442    Display the given Markdown text in a scrollable area using <pre> tag.443 444    Args:445    text (str): The Markdown text to be displayed.446    height (int): Height of the scrollable area in pixels.447    """448    # 使用 <pre> 标签来包裹 Markdown 内容,并添加 CSS 样式创建可滚动的区域449    markdown_container = f"""450    <pre style="451        overflow-y: scroll;452        height: {height}px;453        border: 1px solid #ccc;454        padding: 10px;455        margin-bottom: 20px;456        background-color: #f5f5f5;457    ">458    {text}459    </pre>460    """461 462    st.markdown(markdown_container, unsafe_allow_html=True)463 464 465@st.cache_resource466def compute_one_model_fitted_params(filename, num_fold, query_aware_idxs, resolved_data):467    st.write('---------------')468    one_model_fitted_params = {}469    data = json.load(filename)470    uploaded_labels = [1 if x == "A" else 0 for x in data]471 472    ccount=0473 474    for part in list(query_aware_idxs.keys()):475        if part == "all": continue476        # 使用 st.empty 创建占位符477        progress_text = st.empty()478        # if part not in ["Advice","NLP Tasks"]:continue479        progress_text.write(f"{ccount+1}/{len(list(query_aware_idxs.keys()))-1} "+formal_group_name(part))480        progress_bar = st.progress(0)481        cared_idxs = query_aware_idxs.get(part)482 483        features = []484        labels = []485 486        for idx, item in enumerate(resolved_data):487            if idx not in cared_idxs: continue488            if item['comparison']['accuracy']['comparison'] == 999: continue489            label = uploaded_labels[idx]490            feature = get_feature(item, remove_length=False, way='comparison')491            features.append(feature)492            labels.append(label)493 494        features = np.asarray(features, dtype=np.float32)495        labels = np.asarray(labels)496 497        if num_fold>1:498            np.random.seed(0)499            idxs = np.arange(len(features))500            np.random.shuffle(idxs)501            features = features[idxs]502            labels = labels[idxs]503 504            final_paras = None505            for i in range(num_fold):506                # take the i/10 as test set507                features_len = len(features)508                split_point = int(i / num_fold * features_len)509                features_train, features_test = np.concatenate(510                    [features[:split_point, :], features[split_point + int(features_len / num_fold):, :]],511                    axis=0), features[split_point:split_point + int(features_len / num_fold), :]512                labels_train, labels_test = np.concatenate(513                    [labels[:split_point], labels[split_point + int(features_len / num_fold):]], axis=0), labels[514                                                                                                          split_point:split_point + int(515                                                                                                              features_len / 10)]516                model, parameters = fit_bayes_logistic_regression(features_train, labels_train, scale=0.1)517                if final_paras is None:518                    final_paras = np.asarray(parameters)519                else:520                    final_paras += np.asarray(parameters)521                progress_bar.progress((i + 1)/num_fold)522        else:523            model, parameters = fit_bayes_logistic_regression(features, labels, scale=0.1)524            final_paras = np.asarray(parameters)525            progress_bar.progress(1)526 527        final_paras /= num_fold528        parameters = final_paras.tolist()529        one_model_fitted_params[formal_group_name(part)] = parameters530 531        # 函数处理完毕,清除进度条和文本532        progress_text.empty()533        progress_bar.empty()534        ccount+=1535 536    return one_model_fitted_params537 538def get_json_download_link(json_str, file_name, button_text):539    # 创建一个BytesIO对象540    b64 = base64.b64encode(json_str.encode()).decode()541    href = f'<a href="data:file/json;base64,{b64}" download="{file_name}">{button_text}</a>'542    return href543 544 545if __name__ == "__main__":546    st.title("Visualization of Preference Dissection")547 548    INTRO = """    549This space is used to show visualization results for human and LLM preferences analyzed in the following paper:550 551 552[***Dissecting Human and LLM Preferences***](https://arxiv.org/abs/2402.11296)553 554by [Junlong Li](https://lockon-n.github.io/), [Fan Zhou](https://koalazf99.github.io/), [Shichao Sun](https://shichaosun.github.io/), [Yikai Zhang](https://arist12.github.io/ykzhang/), [Hai Zhao](https://bcmi.sjtu.edu.cn/home/zhaohai/) and [Pengfei Liu](http://www.pfliu.com/)555 556------------557 558Specifically, we include:559 5601. **Complete Preference Dissection in Paper**: shows how the difference of properties in a pair of responses can influence different LLMs'(human included) preference. <br>5612. **Preference Similarity Matrix**: shows the preference similarity among different judges. <br>5623. **Sample-level SHAP Analysis**: applies shapley value to show how the difference of properties in a pair of responses affect the final preference. <br>5634. **Add a New Model for Preference Dissection**: update the preference labels from a new LLM and visualize the results564 565This analysis is based on:566 567> The data we collected here: https://huggingface.co/datasets/GAIR/preference-dissection568 569> The code we released here: https://github.com/GAIR-NLP/Preference-Dissection570"""571    message = custom_md_with_color(INTRO, "DBEFEB")572 573    st.markdown(message, unsafe_allow_html=True)574 575    st.write("## :red[⬇] Click the Box and Select a Section :red[⬇]")576 577    section = st.selectbox(578        "",579        [580            "Complete Preference Dissection in Paper",581            "Preference Similarity Matrix",582            "Sample-level SHAP Analysis",583            'Add a New Model for Preference Dissection'584        ],585    )586    st.markdown("---")587 588    if section == "Complete Preference Dissection in Paper":589        st.header("Complete Preference Dissection in Paper")590        st.markdown("")591        selected_part = st.selectbox(592            "**Scenario/Query-Specific Cases**", list(all_models_fitted_params.keys())593        )594 595        models = all_models_fitted_params[selected_part]596 597        model_names = list(models.keys())598        selected_models = st.multiselect(599            "**Select LLMs (Human) to display**",600            model_names,601            default=["human", "gpt-4-1106-preview"],602        )603 604        st.text(605            "The value for each property indicates that, when response A satisfies only this\nproperty better than response B and all else equal, the probability of response\nA being preferred.")606 607        if len(selected_models) > 0:608            show_all_models_prob(models, selected_models, feature_names_short)609        else:610            st.write("Please select at least one model to display.")611    elif section == "Preference Similarity Matrix":612        st.header("Preference Similarity Matrix")613 614        # Initialize session state for similarity matrix615 616        # convert `groupwise_fitted_paras` to `modelwise_fitted_paras`617 618        models = list(modelwise_fitted_paras.keys())619        # Option to choose between preset models or selecting models620        option = st.radio(621            "**Choose your models setting**",622            ("Use Preset Models", "Select Models Manually"),623        )624 625        if option == "Use Preset Models":626            selected_models = pre_set_full_model_order627        else:628            selected_models = st.multiselect(629                "**Select Models**", models, default=models[:5]630            )631 632        # Input for threshold value633        st.text(634            "The similarity bewteen two judges is the average pearson correlation coefficient of\nthe fitted Bayesian logistic regression models' weights across all scenarios.")635 636        selected_part = st.selectbox(637            "**Overall or Scenario/Query-Specific Cases**", ["Overall"] + list(all_models_fitted_params.keys())638        )639 640        st.text(641            "\"Overall\" is the average similarity across all scenarios, \nwhile \"Scenario/Query-Specific Cases\" is the similarity within \nthe selected scenario/query-specific cases.")642 643        if len(selected_models) >= 2:644            # Call the cached function645            similarity_matrix = calculate_similarity_matrix(646                modelwise_fitted_paras, selected_models, "pearson", selected_part647            )648            # Store the matrix in session state649            # Slider to adjust figure size650            fig_size = (651                25652                if option == "Use Preset Models"653                else int(33 * len(selected_models) / 25)654            )655 656            plt.figure(figsize=(fig_size * 1.1, fig_size))657            ax = sns.heatmap(658                similarity_matrix,659                annot=True,660                annot_kws={"size": 18},  # Change annotation font size661                xticklabels=[become_formal(x) for x in selected_models],662                yticklabels=[become_formal(x) for x in selected_models],663            )664 665            # Add this line to get the colorbar object666            cbar = ax.collections[0].colorbar667 668            # Here, specify the font size for the colorbar669            for label in cbar.ax.get_yticklabels():670                # label.set_fontsize(20)  # Set the font size (change '10' as needed)671                label.set_fontname(672                    "Times New Roman"673                )  # Set the font name (change as needed)674 675            plt.xticks(rotation=45, fontname="Times New Roman", ha="right")676            plt.yticks(rotation=0, fontname="Times New Roman")677 678            plt.tight_layout()679            st.pyplot(plt)680        else:681            st.warning("Please select at least two models.")682    elif section == "Sample-level SHAP Analysis":683        st.header("Sample-level SHAP Analysis")684        resolved_data_file = "./data/chatbot_arena_no-tie_group_balanced_resolved.jsonl"685        source_data_file = "./data/chatbot_arena_shuffled_no-tie_group_balanced.jsonl"686        reference_data_file = (687            "./data/chatbot_arena_shuffled_no-tie_gpt4_ref_group_balanced.jsonl"688        )689 690        # Load and prepare data691        resolved_data, source_data, reference_data = (692            read_all(resolved_data_file),693            read_all(source_data_file),694            read_all(reference_data_file),695        )696        ok_idxs = [697            i698            for i, item in enumerate(resolved_data)699            if item["comparison"]["accuracy"]["comparison"] != 999700        ]701        resolved_data, source_data, reference_data = (702            [resolved_data[i] for i in ok_idxs],703            [source_data[i] for i in ok_idxs],704            [reference_data[i] for i in ok_idxs],705        )706        features = np.asarray(707            [708                get_feature(item, remove_length=False, way="comparison")709                for item in resolved_data710            ],711            dtype=np.float32,712        )713 714        # Initialize the index715        if "sample_ind" not in st.session_state:716            st.session_state.sample_ind = 0717 718 719        # Function to update the index720        def update_index(change):721            st.session_state.sample_ind += change722            st.session_state.sample_ind = max(723                0, min(st.session_state.sample_ind, len(features) - 1)724            )725 726 727        col1, col2, col3, col4, col5 = st.columns([1, 2, 1, 2, 1])728 729        with col1:730            st.button("Prev", on_click=update_index, args=(-1,))731 732        with col3:733            number = st.number_input(734                "Go to sample:",735                min_value=0,736                max_value=len(features) - 1,737                value=st.session_state.sample_ind,738            )739            if number != st.session_state.sample_ind:740                st.session_state.sample_ind = number741 742        with col5:743            st.button("Next", on_click=update_index, args=(1,))744 745        # Use the updated sample index746        sample_ind = st.session_state.sample_ind747 748        reference, source, resolved = (749            reference_data[sample_ind],750            source_data[sample_ind],751            resolved_data[sample_ind],752        )753 754        groups = [f"[Scenario] {source['group']}"] + process_query_info(755            resolved["query_info"]756        )757 758        st.write("")759        group = st.selectbox(760            "**Scenario & Potential Query-Specific Cases:**\n\nWe set the scenario of this sample by default, but you can also select certain query-specfic groups if the query satisfy certain conditions.",761            options=groups,762        )763        model_name = st.selectbox(764            "**The Preference of which LLM (Human):**",765            options=list(all_models_fitted_params[group].keys()),766        )767        paras_spec = all_models_fitted_params[group][model_name]768        model = BayesianLogisticRegression(paras_spec)769        explainer = shap.Explainer(model=model.return_prob, masker=np.zeros((1, 29)))770 771        # Calculate SHAP values772        shap_values = explainer(773            features[st.session_state.sample_ind: st.session_state.sample_ind + 1, :]774        )775        shap_values.feature_names = list(feature_name_to_id_short.keys())776 777        # Plotting778 779        st.markdown(780            "> *f(x) > 0.5 means response A is preferred more, and vice versa.*"781        )782        st.markdown(783            "> *Property = 1 means response A satisfy the property better than B, and vice versa. We only show the properties that distinguish A and B.*"784        )785 786        # count how mant nonzero in shape_values[0].data787        nonzero = np.nonzero(shap_values[0].data)[0].shape[0]788        shap.plots.waterfall(shap_values[0], max_display=nonzero + 1, show=False)789        fig = plt.gcf()790        st.pyplot(fig)791 792        # st.subheader(793        #     "**Detailed information (source data and annotation) of this sample.**"794        # )795 796        # We pop some attributes first797 798        # RAW Json799        simplified_source = {800            "query": source["prompt"],801            f"response A ({source['model_a']}, {source['response_a word']} words)": source[802                "response_a"803            ],804            f"response B ({source['model_b']}, {source['response_b word']} words)": source[805                "response_b"806            ],807            "GPT-4-Turbo Reference": reference["output"],808        }809        simplified_resolved = {810            "query-specific:": resolved["query_info"],811            "Annotation": {812                k: v["meta"]813                for k, v in resolved["comparison"].items()814                if v["meta"] is not None and k != "length"815            },816        }817 818        # Source Data Rendering819        # st.json(simplified_source)820        st.write("#### Source Data")821        st.text_area(822            "**Query**:\n",823            f"""{source["prompt"]}\n""",824        )825        st.text_area(826            f"**response A ({source['model_a']}, {source['response_a word']} words)**:\n",827            f"""{source["response_a"]}\n""",828            height=200,829        )830        st.text_area(831            f"**response B ({source['model_b']}, {source['response_b word']} words)**:\n",832            f"""{source["response_b"]}\n""",833            height=200,834        )835        st.text_area(836            f"**GPT-4-Turbo Reference**:\n",837            f"""{reference["output"]}\n""",838            height=200,839        )840 841        # Resolved Data Rendering842        st.markdown("---")843        st.write("### Annotation")844        # st.json(simplified_resolved)845        st.write("#### Query Information\n")846        query_info = json_to_markdown_bold_keys(simplified_resolved["query-specific:"])847        st.markdown(custom_md_with_color(query_info, "DFEFDB"), unsafe_allow_html=True)848 849        specific_check_feature_fixed = [850            "length",851            "accuracy",852        ]853        specific_check_feature_dynamic = [854            "clarify user intent",855            "showing empathetic",856            "satisfying explicit constraints",857            "supporting explicit subjective stances",858            "correcting explicit mistakes or biases"859        ]860        specific_check_feature = specific_check_feature_fixed + specific_check_feature_dynamic861        normal_check_feature = {862            k: v["meta"]863            for k, v in resolved["comparison"].items()864            if v["meta"] is not None and k not in specific_check_feature865        }866        # generate table for normal check feature867        data = {"Category": [], "Response 1": [], "Response 2": []}868 869        for category, responses in normal_check_feature.items():870            # print(responses)871            data["Category"].append(category)872            data["Response 1"].append(responses["Response 1"])873            data["Response 2"].append(responses["Response 2"])874 875        df = pd.DataFrame(data)876 877        # Display the table in Streamlit878        st.write("#### Ratings of Basic Properties\n")879        st.table(df)880 881        # specific check features: 'accuracy', and 'satisfying explicit constraints'882        st.write("#### Error Detection")883 884        # xx885        acc1 = simplified_resolved["Annotation"]["accuracy"]["Response 1"]886        newacc1 = {"applicable to detect errors": acc1["accuracy check"],887                   "detected errors": acc1["inaccuracies"]}888        acc2 = simplified_resolved["Annotation"]["accuracy"]["Response 2"]889        newacc2 = {"applicable to detect errors": acc2["accuracy check"],890                   "detected errors": acc2["inaccuracies"]}891 892        # Convert the JSON to a Markdown string893        response_1 = json_to_markdown_bold_keys(newacc1)894        response_2 = json_to_markdown_bold_keys(newacc2)895        st.markdown("##### Response 1")896        st.markdown(custom_md_with_color(response_1, "DBE7EF"), unsafe_allow_html=True)897        st.text("")898        st.markdown("##### Response 2")899        st.markdown(custom_md_with_color(response_2, "DBE7EF"), unsafe_allow_html=True)900 901        if any(j in simplified_resolved['Annotation'] for j in specific_check_feature_dynamic):902            st.text("")903            st.markdown("#### Query-Specific Annotation")904 905            for j in specific_check_feature_dynamic:906                if j in simplified_resolved['Annotation']:907                    st.write(f"**{j} (ratings from 0-3 or specific labels)**")908                    st.markdown(custom_md_with_color(json_to_markdown_bold_keys(simplified_resolved['Annotation'][j]),909                                                     "E8DAEF"), unsafe_allow_html=True)910                    st.text("")911    else:912        st.header("Add a New Model for Preference Dissection")913        resolved_data = read_all("./data/chatbot_arena_no-tie_group_balanced_resolved.jsonl")914        query_aware_idxs = read_all("./data/query_aware_idxs.json")915 916        st.write("Upload the preference labels from a new LLM.")917        st.write("The data in ths .json file should be a list with 5240 (the same as the data size) elements, each belongs to {\"A\",\"B\"} indicating the preferred one in each pair.")918        st.write("We provide an example in ```./data/example_preference_labels.json``` in the ``Files`` of the space, which are the preference labels of human.")919        filename = st.file_uploader("", type=["json"],920                                    key="new_model_fitted_params")921 922        one_model_fitted_params = None923 924        if filename is not None:925            st.write("Uploaded successfully.")926 927            st.write("Please select the number of folds for fitting the models. 1 means no multi-fold averaging. (Warning! Large number of fold may cause OOM and the crush of this space.)")928            num_fold = st.selectbox("Number of Folds", [1, 2, 5, 10], index=0)929 930            one_model_fitted_params = compute_one_model_fitted_params(filename, num_fold, query_aware_idxs,931                                                                      resolved_data)932 933        if one_model_fitted_params is not None:934            json_data = json.dumps(one_model_fitted_params, indent=4)935            st.markdown(get_json_download_link(json_data, "fitted_weights.json", "Download Fitted Bayesian Logistic Models Weights"), unsafe_allow_html=True)936 937            st.write("The visualization is the same as the first section.")938 939            selected_part = st.selectbox("**Scenario/Query-Specific Cases**", list(one_model_fitted_params.keys()))940            weights = one_model_fitted_params[selected_part]941            show_one_model_prob(weights, feature_names_short)