GAIR/Preference-Dissection-Visualization
6
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)