import plotly.express as px
import plotly.graph_objects as go
import pandas as pd
METRICS_FOR_PLOTS = ["Accuracy", "MCC", "Weighted F1"]
# Color scheme matching the UI theme
COLORS = {
'primary': '#667eea',
'secondary': '#764ba2',
'accent': '#F97316',
'gradient': ['#667eea', '#764ba2', '#F97316', '#06b6d4', '#10b981']
}
def extract_mean_std(col):
if isinstance(col, str) and '±' in col:
try:
mean, std = col.split('±')
return float(mean.strip()), float(std.strip())
except:
return None, None
elif isinstance(col, (int, float)):
return col, 0
return None, None
def prepare_leaderboard_df(df):
for metric in METRICS_FOR_PLOTS:
means, stds = zip(*df[metric].apply(extract_mean_std))
df[f"{metric}_mean"] = means
df[f"{metric}_std"] = stds
return df
def make_plot_wrapper(leaderboard_df, group_by, filter_col, orientation):
def plot_fn(filter_val, metric, dataset):
return plot_metric_bar(
df=leaderboard_df,
group_by=group_by,
filter_col=filter_col,
dataset=str(dataset),
filter_val=filter_val,
metric=metric,
orientation=orientation
)
return plot_fn
def make_model_all_datasets_wrapper(leaderboard_df):
def plot_fn(model, metric):
return plot_metric_bar_all_datasets(
df=leaderboard_df,
model=model,
metric=metric
)
return plot_fn
def plot_metric_bar_all_datasets(df, model, metric, orientation="v"):
df = df.copy()
df = df[df['Model'] == model]
y_col = f"{metric}_mean"
std_col = f"{metric}_std"
if df.empty or y_col not in df.columns:
return px.bar(title=f"No data found for model: {model}")
df = df.sort_values(by=y_col, ascending=False)
x_axis, y_axis = (y_col, "Task") if orientation == "h" else ("Task", y_col)
height = 20 * len(df) + 200 if orientation == "h" else None
fig = px.bar(
df,
x=x_axis,
y=y_axis,
color='Dataset Name',
orientation=orientation,
text_auto=".2f",
labels={y_col: metric, "Task": "Task"},
title=f"{metric} for {model} across all datasets",
height=height,
color_discrete_sequence=COLORS['gradient']
)
if orientation == "h":
hovertemplate = (
"%{y}
"
f"Mean {metric}: %{{x:.2f}}
"
"Std Dev: %{customdata[0]:.3f}"
""
)
else:
hovertemplate = (
"%{x}
"
f"Mean {metric}: %{{y:.2f}}
"
"Std Dev: %{customdata[0]:.3f}"
""
)
fig.update_traces(
hovertemplate=hovertemplate,
customdata=df[[std_col]].values
)
layout_args = dict(
title={
'text': f"{metric} for {model} across all tasks",
'x': 0.5,
'xanchor': 'center',
'yanchor': 'top',
'font': dict(size=22, color='#2d3748', family='Inter, sans-serif'),
'pad': dict(t=20, b=10),
},
margin=dict(t=80, b=100, l=150, r=10),
plot_bgcolor='rgba(0,0,0,0)',
paper_bgcolor='white',
font=dict(family='Inter, sans-serif', color='#2d3748'),
)
if orientation == "h":
layout_args["xaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
layout_args["yaxis_tickfont_size"] = 10
layout_args["yaxis"] = dict(gridcolor='#e2e8f0')
else:
layout_args["yaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
layout_args["xaxis"] = dict(gridcolor='#e2e8f0')
fig.update_layout(**layout_args)
return fig
def plot_metric_bar(df, group_by, filter_col, filter_val, dataset, metric, orientation="v"):
df = df.copy()
df = df[df['Dataset Name'] == dataset]
df = df[df[filter_col] == filter_val]
y_col = f"{metric}_mean"
std_col = f"{metric}_std"
if df.empty or y_col not in df.columns:
return px.bar(title="No data found.")
df = df.sort_values(by=y_col, ascending=False)
x_axis, y_axis = (y_col, group_by) if orientation == "h" else (group_by, y_col)
height = 20 * len(df) + 200 if orientation == "h" else None
fig = px.bar(
df,
x=x_axis,
y=y_axis,
orientation=orientation,
text_auto=".2f",
labels={y_col: metric, group_by: group_by},
title=f"{metric} for {filter_val} - {dataset}",
height=height,
)
if orientation == "h":
hovertemplate = (
"%{y}
"
f"Mean {metric}: %{{x:.2f}}
"
"Std Dev: %{customdata[0]:.3f}"
""
)
else:
hovertemplate = (
"%{x}
"
f"Mean {metric}: %{{y:.2f}}
"
"Std Dev: %{customdata[0]:.3f}"
""
)
fig.update_traces(
marker_color=COLORS['primary'],
marker_line_color=COLORS['secondary'],
marker_line_width=1.5,
hovertemplate=hovertemplate,
customdata=df[[std_col]].values
)
layout_args = dict(
title={
'text': f"{metric} for {filter_val} - {dataset}",
'x': 0.5,
'xanchor': 'center',
'yanchor': 'top',
'font': dict(size=22, color='#2d3748', family='Inter, sans-serif'),
'pad': dict(t=20, b=10),
},
margin=dict(t=80, b=100, l=150, r=10),
plot_bgcolor='rgba(0,0,0,0)',
paper_bgcolor='white',
font=dict(family='Inter, sans-serif', color='#2d3748'),
)
if orientation == "h":
layout_args["xaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
layout_args["yaxis_tickfont_size"] = 10
layout_args["yaxis"] = dict(gridcolor='#e2e8f0')
else:
layout_args["yaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
layout_args["xaxis"] = dict(gridcolor='#e2e8f0')
fig.update_layout(**layout_args)
return fig