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