dna-benchmark / app.py
pefidias's picture
UI Revamp (#6)
9f99eeb verified
Raw History Blame
7.4 kB
import gradio as gr
from apscheduler.schedulers.background import BackgroundScheduler
from gradio_leaderboard import ColumnFilter, Leaderboard, SelectColumns
from huggingface_hub import snapshot_download
from src.about import INTRODUCTION_TEXT, LLM_BENCHMARKS_TEXT, TITLE
from src.display.css_html_js import custom_css
from src.display.utils import BenchRawColumn, fields
from src.envs import API, EVAL_RESULTS_PATH, REPO_ID, RESULTS_REPO, TOKEN
from src.populate import get_leaderboard_df_from_hf_dataset, summarize_model_task_type_performance
from src.display.plotting import METRICS_FOR_PLOTS, extract_mean_std, prepare_leaderboard_df, make_plot_wrapper, make_model_all_datasets_wrapper, plot_metric_bar
def restart_space():
API.restart_space(repo_id=REPO_ID)
try:
print(EVAL_RESULTS_PATH)
snapshot_download(
repo_id=RESULTS_REPO,
local_dir=EVAL_RESULTS_PATH,
repo_type="dataset",
tqdm_class=None,
etag_timeout=30,
token=TOKEN,
)
LEADERBOARD_DF = get_leaderboard_df_from_hf_dataset(EVAL_RESULTS_PATH)
LEADERBOARD_DF = summarize_model_task_type_performance(LEADERBOARD_DF)
LEADERBOARD_DF = prepare_leaderboard_df(LEADERBOARD_DF)
wrapped_task_plot = make_plot_wrapper(leaderboard_df=LEADERBOARD_DF, group_by="Model", filter_col="Task", orientation="h")
wrapped_model_all_datasets_plot = make_model_all_datasets_wrapper(leaderboard_df=LEADERBOARD_DF)
except Exception as e:
print(e)
restart_space()
def init_leaderboard(dataframe):
if dataframe is None or dataframe.empty:
raise ValueError("Leaderboard DataFrame is empty or None.")
return Leaderboard(
value=dataframe,
datatype=[c.type for c in fields(BenchRawColumn)],
select_columns=SelectColumns(
default_selection=[c.name for c in fields(BenchRawColumn) if c.displayed_by_default],
cant_deselect=[c.name for c in fields(BenchRawColumn) if c.never_hidden],
label="Select Columns to Display:",
),
filter_columns=[
ColumnFilter(BenchRawColumn.task.name, type="checkboxgroup", label=BenchRawColumn.task.name),
# ColumnFilter(BenchRawColumn.precision.name, type="checkboxgroup", label="Precision"),
ColumnFilter(
BenchRawColumn.model_params.name,
type="slider",
min=100,
max=10000,
label="Select the number of parameters (M)",
),
ColumnFilter(
BenchRawColumn.embds_dim.name,
type="slider",
min=10,
max=10000,
label="Select the Embeddings Size",
),
# ColumnFilter(BenchRawColumn.still_on_hub.name, type="boolean", label="Deleted/incomplete", default=True),
ColumnFilter(
BenchRawColumn.max_context_len.name,
type="slider",
min=10,
max=10000000,
label="Select the Max Context Length (bp)",
),
],
search_columns=[BenchRawColumn.model.name],
)
# Create an elegant dark theme with custom colors for checkboxes and sliders
theme = gr.themes.Soft(
primary_hue="violet",
secondary_hue="purple",
neutral_hue="slate",
spacing_size="md",
radius_size="lg",
text_size="md",
font=[gr.themes.GoogleFont("Inter"), "ui-sans-serif", "system-ui", "sans-serif"],
).set(
checkbox_background_color_selected="*primary_600",
checkbox_border_color_selected="*primary_600",
checkbox_background_color_selected_dark="*primary_500",
checkbox_border_color_selected_dark="*primary_500",
slider_color="*primary_600",
slider_color_dark="*primary_500",
button_primary_background_fill="*primary_600",
button_primary_background_fill_hover="*primary_700",
button_primary_background_fill_dark="*primary_500",
button_primary_background_fill_hover_dark="*primary_600",
)
demo = gr.Blocks(
css=custom_css,
theme=theme,
js="""
() => {
const theme = localStorage.getItem('theme');
if (theme === null) {
localStorage.setItem('theme', 'dark');
document.body.classList.add('dark');
}
}
""",
title="DNA Benchmark",
head="<link rel='icon' href='data:image/svg+xml,<svg xmlns=\"http://www.w3.org/2000/svg\" viewBox=\"0 0 100 100\"><text y=\".9em\" font-size=\"90\">🧬</text></svg>'>"
)
with demo:
gr.HTML(TITLE)
gr.Markdown(INTRODUCTION_TEXT, elem_classes="markdown-text")
with gr.Tabs(elem_classes="tab-buttons") as tabs:
with gr.TabItem("🏅 LLM Leaderboard", elem_id="llm-benchmark-tab-table", id=0):
leaderboard = init_leaderboard(LEADERBOARD_DF)
with gr.TabItem("📊 Performance Plots", elem_id="llm-benchmark-performance-plots", id=1):
with gr.Accordion("🧬 Overview of Performance for a Specific Task", open=False):
gr.Markdown("Visualize model performance for a specific task and metric")
task_choices = sorted(LEADERBOARD_DF["Task"].dropna().unique())
dataset_choices = sorted(LEADERBOARD_DF["Dataset Name"].dropna().unique())
with gr.Row():
task_dropdown = gr.Dropdown(choices=task_choices, label="Select Task")
metric_dropdown = gr.Dropdown(choices=METRICS_FOR_PLOTS, value="Accuracy", label="Select Metric")
dataset_dropdown = gr.Dropdown(choices=dataset_choices, value="InstaDeepAI/nucleotide_transformer_downstream_tasks_revised", label="Select Dataset")
performance_plot = gr.Plot()
task_dropdown.change(fn=wrapped_task_plot, inputs=[task_dropdown, metric_dropdown, dataset_dropdown], outputs=performance_plot)
metric_dropdown.change(fn=wrapped_task_plot, inputs=[task_dropdown, metric_dropdown, dataset_dropdown], outputs=performance_plot)
dataset_dropdown.change(fn=wrapped_task_plot, inputs=[task_dropdown, metric_dropdown, dataset_dropdown], outputs=performance_plot)
with gr.Accordion("🤖 Overview of Performance for a Specific Model", open=False):
gr.Markdown("Visualize model performance across all tasks from all datasets")
model_choices = sorted(LEADERBOARD_DF["Model"].dropna().unique())
with gr.Row():
model_dropdown = gr.Dropdown(choices=model_choices, label="Select Model")
metric2_dropdown = gr.Dropdown(choices=METRICS_FOR_PLOTS, value="Accuracy", label="Select Metric")
model_plot = gr.Plot()
model_dropdown.change(fn=wrapped_model_all_datasets_plot, inputs=[model_dropdown, metric2_dropdown], outputs=model_plot)
metric2_dropdown.change(fn=wrapped_model_all_datasets_plot, inputs=[model_dropdown, metric2_dropdown], outputs=model_plot)
with gr.TabItem("📝 About", elem_id="llm-benchmark-about", id=2):
gr.Markdown(LLM_BENCHMARKS_TEXT, elem_classes="markdown-text")
scheduler = BackgroundScheduler()
scheduler.add_job(restart_space, "interval", seconds=1800)
scheduler.start()
demo.queue(default_concurrency_limit=40).launch(share=True)