CoolFace
Apppublic

lmsys/mt-bench

sourceHugging Faceotherupdated 3y agoView on Hugging Face
203likes
app.py430 linesDownload Raw Back to root
1"""2Usage:3python3 qa_browser.py --share4"""5 6import argparse7from collections import defaultdict8import re9 10import gradio as gr11 12from common import (13    load_questions,14    load_model_answers,15    load_single_model_judgments,16    load_pairwise_model_judgments,17    resolve_single_judgment_dict,18    resolve_pairwise_judgment_dict,19    get_single_judge_explanation,20    get_pairwise_judge_explanation,21)22 23 24questions = []25model_answers = {}26 27model_judgments_normal_single = {}28model_judgments_math_single = {}29 30model_judgments_normal_pairwise = {}31model_judgments_math_pairwise = {}32 33question_selector_map = {}34category_selector_map = defaultdict(list)35 36 37def display_question(category_selector, request: gr.Request):38    choices = category_selector_map[category_selector]39    return gr.Dropdown.update(40        value=choices[0],41        choices=choices,42    )43 44 45def display_pairwise_answer(46    question_selector, model_selector1, model_selector2, request: gr.Request47):48    q = question_selector_map[question_selector]49    qid = q["question_id"]50 51    ans1 = model_answers[model_selector1][qid]52    ans2 = model_answers[model_selector2][qid]53 54    chat_mds = pairwise_to_gradio_chat_mds(q, ans1, ans2)55    gamekey = (qid, model_selector1, model_selector2)56 57    judgment_dict = resolve_pairwise_judgment_dict(58        q,59        model_judgments_normal_pairwise,60        model_judgments_math_pairwise,61        multi_turn=False,62    )63 64    explanation = (65        "##### Model Judgment (first turn)\n"66        + get_pairwise_judge_explanation(gamekey, judgment_dict)67    )68 69    judgment_dict_turn2 = resolve_pairwise_judgment_dict(70        q,71        model_judgments_normal_pairwise,72        model_judgments_math_pairwise,73        multi_turn=True,74    )75 76    explanation_turn2 = (77        "##### Model Judgment (second turn)\n"78        + get_pairwise_judge_explanation(gamekey, judgment_dict_turn2)79    )80 81    return chat_mds + [explanation] + [explanation_turn2]82 83 84def display_single_answer(question_selector, model_selector1, request: gr.Request):85    q = question_selector_map[question_selector]86    qid = q["question_id"]87 88    ans1 = model_answers[model_selector1][qid]89 90    chat_mds = single_to_gradio_chat_mds(q, ans1)91    gamekey = (qid, model_selector1)92 93    judgment_dict = resolve_single_judgment_dict(94        q, model_judgments_normal_single, model_judgments_math_single, multi_turn=False95    )96 97    explanation = "##### Model Judgment (first turn)\n" + get_single_judge_explanation(98        gamekey, judgment_dict99    )100 101    judgment_dict_turn2 = resolve_single_judgment_dict(102        q, model_judgments_normal_single, model_judgments_math_single, multi_turn=True103    )104 105    explanation_turn2 = (106        "##### Model Judgment (second turn)\n"107        + get_single_judge_explanation(gamekey, judgment_dict_turn2)108    )109 110    return chat_mds + [explanation] + [explanation_turn2]111 112 113newline_pattern1 = re.compile("\n\n(\d+\. )")114newline_pattern2 = re.compile("\n\n(- )")115 116 117def post_process_answer(x):118    """Fix Markdown rendering problems."""119    x = x.replace("\u2022", "- ")120    x = re.sub(newline_pattern1, "\n\g<1>", x)121    x = re.sub(newline_pattern2, "\n\g<1>", x)122    return x123 124 125def pairwise_to_gradio_chat_mds(question, ans_a, ans_b, turn=None):126    end = len(question["turns"]) if turn is None else turn + 1127 128    mds = ["", "", "", "", "", "", ""]129    for i in range(end):130        base = i * 3131        if i == 0:132            mds[base + 0] = "##### User\n" + question["turns"][i]133        else:134            mds[base + 0] = "##### User's follow-up question \n" + question["turns"][i]135        mds[base + 1] = "##### Assistant A\n" + post_process_answer(136            ans_a["choices"][0]["turns"][i].strip()137        )138        mds[base + 2] = "##### Assistant B\n" + post_process_answer(139            ans_b["choices"][0]["turns"][i].strip()140        )141 142    ref = question.get("reference", ["", ""])143 144    ref_md = ""145    if turn is None:146        if ref[0] != "" or ref[1] != "":147            mds[6] = f"##### Reference Solution\nQ1. {ref[0]}\nQ2. {ref[1]}"148    else:149        x = ref[turn] if turn < len(ref) else ""150        if x:151            mds[6] = f"##### Reference Solution\n{ref[turn]}"152        else:153            mds[6] = ""154    return mds155 156 157def single_to_gradio_chat_mds(question, ans, turn=None):158    end = len(question["turns"]) if turn is None else turn + 1159 160    mds = ["", "", "", "", ""]161    for i in range(end):162        base = i * 2163        if i == 0:164            mds[base + 0] = "##### User\n" + question["turns"][i]165        else:166            mds[base + 0] = "##### User's follow-up question \n" + question["turns"][i]167        mds[base + 1] = "##### Assistant A\n" + post_process_answer(168            ans["choices"][0]["turns"][i].strip()169        )170 171    ref = question.get("reference", ["", ""])172 173    ref_md = ""174    if turn is None:175        if ref[0] != "" or ref[1] != "":176            mds[4] = f"##### Reference Solution\nQ1. {ref[0]}\nQ2. {ref[1]}"177    else:178        x = ref[turn] if turn < len(ref) else ""179        if x:180            mds[4] = f"##### Reference Solution\n{ref[turn]}"181        else:182            mds[4] = ""183    return mds184 185 186def build_question_selector_map():187    global question_selector_map, category_selector_map188 189    # Build question selector map190    for q in questions:191        preview = f"{q['question_id']}: " + q["turns"][0][:128] + "..."192        question_selector_map[preview] = q193        category_selector_map[q["category"]].append(preview)194 195 196def sort_models(models):197    priority = {198        "Llama-2-70b-chat": "aaaa", 199        "Llama-2-13b-chat": "aaab",200        "Llama-2-7b-chat": "aaac",201    }202 203    models = list(models)204    models.sort(key=lambda x: priority.get(x, x))205    return models206 207 208def build_pairwise_browser_tab():209    global question_selector_map, category_selector_map210 211    models = sort_models(list(model_answers.keys()))212    num_sides = 2213    num_turns = 2214    side_names = ["A", "B"]215 216    question_selector_choices = list(question_selector_map.keys())217    category_selector_choices = list(category_selector_map.keys())218 219    # Selectors220    with gr.Row():221        with gr.Column(scale=1, min_width=200):222            category_selector = gr.Dropdown(223                choices=category_selector_choices, label="Category", container=False224            )225        with gr.Column(scale=100):226            question_selector = gr.Dropdown(227                choices=question_selector_choices, label="Question", container=False228            )229 230    model_selectors = [None] * num_sides231    with gr.Row():232        for i in range(num_sides):233            with gr.Column():234                if i == 0:235                    value = models[0]236                else:237                    value = "gpt-3.5-turbo"238                model_selectors[i] = gr.Dropdown(239                    choices=models,240                    value=value,241                    label=f"Model {side_names[i]}",242                    container=False,243                )244 245    # Conversation246    chat_mds = []247    for i in range(num_turns):248        chat_mds.append(gr.Markdown(elem_id=f"user_question_{i+1}"))249        with gr.Row():250            for j in range(num_sides):251                with gr.Column(scale=100):252                    chat_mds.append(gr.Markdown())253 254                if j == 0:255                    with gr.Column(scale=1, min_width=8):256                        gr.Markdown()257    reference = gr.Markdown(elem_id=f"reference")258    chat_mds.append(reference)259 260    model_explanation = gr.Markdown(elem_id="model_explanation")261    model_explanation2 = gr.Markdown(elem_id="model_explanation")262 263    # Callbacks264    category_selector.change(display_question, [category_selector], [question_selector])265    question_selector.change(266        display_pairwise_answer,267        [question_selector] + model_selectors,268        chat_mds + [model_explanation] + [model_explanation2],269    )270 271    for i in range(num_sides):272        model_selectors[i].change(273            display_pairwise_answer,274            [question_selector] + model_selectors,275            chat_mds + [model_explanation] + [model_explanation2],276        )277 278    return (category_selector,)279 280 281def build_single_answer_browser_tab():282    global question_selector_map, category_selector_map283 284    models = sort_models(list(model_answers.keys()))285    num_sides = 1286    num_turns = 2287    side_names = ["A"]288 289    question_selector_choices = list(question_selector_map.keys())290    category_selector_choices = list(category_selector_map.keys())291 292    # Selectors293    with gr.Row():294        with gr.Column(scale=1, min_width=200):295            category_selector = gr.Dropdown(296                choices=category_selector_choices, label="Category", container=False297            )298        with gr.Column(scale=100):299            question_selector = gr.Dropdown(300                choices=question_selector_choices, label="Question", container=False301            )302 303    model_selectors = [None] * num_sides304    with gr.Row():305        for i in range(num_sides):306            with gr.Column():307                model_selectors[i] = gr.Dropdown(308                    choices=models,309                    value=models[i] if len(models) > i else "",310                    label=f"Model {side_names[i]}",311                    container=False,312                )313 314    # Conversation315    chat_mds = []316    for i in range(num_turns):317        chat_mds.append(gr.Markdown(elem_id=f"user_question_{i+1}"))318        with gr.Row():319            for j in range(num_sides):320                with gr.Column(scale=100):321                    chat_mds.append(gr.Markdown())322 323                if j == 0:324                    with gr.Column(scale=1, min_width=8):325                        gr.Markdown()326 327    reference = gr.Markdown(elem_id=f"reference")328    chat_mds.append(reference)329 330    model_explanation = gr.Markdown(elem_id="model_explanation")331    model_explanation2 = gr.Markdown(elem_id="model_explanation")332 333    # Callbacks334    category_selector.change(display_question, [category_selector], [question_selector])335    question_selector.change(336        display_single_answer,337        [question_selector] + model_selectors,338        chat_mds + [model_explanation] + [model_explanation2],339    )340 341    for i in range(num_sides):342        model_selectors[i].change(343            display_single_answer,344            [question_selector] + model_selectors,345            chat_mds + [model_explanation] + [model_explanation2],346        )347 348    return (category_selector,)349 350 351block_css = """352#user_question_1 {353    background-color: #DEEBF7;354}355#user_question_2 {356    background-color: #E2F0D9;357}358#reference {359    background-color: #FFF2CC;360}361#model_explanation {362    background-color: #FBE5D6;363}364"""365 366 367def load_demo():368    dropdown_update = gr.Dropdown.update(value=list(category_selector_map.keys())[0])369    return dropdown_update, dropdown_update370 371 372def build_demo():373    build_question_selector_map()374 375    with gr.Blocks(376        title="MT-Bench Browser",377        theme=gr.themes.Base(text_size=gr.themes.sizes.text_lg),378        css=block_css,379    ) as demo:380        gr.Markdown(381            """382# MT-Bench Browser383| [Paper](https://arxiv.org/abs/2306.05685) | [Code](https://github.com/lm-sys/FastChat/tree/main/fastchat/llm_judge) | [Leaderboard](https://huggingface.co/spaces/lmsys/chatbot-arena-leaderboard) |384"""385        )386        with gr.Tab("Single Answer Grading"):387            (category_selector,) = build_single_answer_browser_tab()388        with gr.Tab("Pairwise Comparison"):389            (category_selector2,) = build_pairwise_browser_tab()390        demo.load(load_demo, [], [category_selector, category_selector2])391 392    return demo393 394 395if __name__ == "__main__":396    parser = argparse.ArgumentParser()397    parser.add_argument("--host", type=str, default="0.0.0.0")398    parser.add_argument("--port", type=int)399    parser.add_argument("--share", action="store_true")400    parser.add_argument("--bench-name", type=str, default="mt_bench")401    args = parser.parse_args()402    print(args)403 404    question_file = f"data/{args.bench_name}/question.jsonl"405    answer_dir = f"data/{args.bench_name}/model_answer"406    pairwise_model_judgment_file = (407        f"data/{args.bench_name}/model_judgment/gpt-4_pair.jsonl"408    )409    single_model_judgment_file = (410        f"data/{args.bench_name}/model_judgment/gpt-4_single.jsonl"411    )412 413    # Load questions414    questions = load_questions(question_file, None, None)415 416    # Load answers417    model_answers = load_model_answers(answer_dir)418 419    # Load model judgments420    model_judgments_normal_single = (421        model_judgments_math_single422    ) = load_single_model_judgments(single_model_judgment_file)423    model_judgments_normal_pairwise = (424        model_judgments_math_pairwise425    ) = load_pairwise_model_judgments(pairwise_model_judgment_file)426 427    demo = build_demo()428    demo.launch(429        server_name=args.host, server_port=args.port, share=args.share, max_threads=200430    )