Restore grouped comparison chart
Browse files
app.py
CHANGED
|
@@ -982,6 +982,36 @@ def metric_chart(results: pd.DataFrame, selected_metrics: list[str] | None = Non
|
|
| 982 |
return fig
|
| 983 |
|
| 984 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 985 |
def time_chart(results: pd.DataFrame) -> go.Figure:
|
| 986 |
if results is None or results.empty or "seconds" not in results:
|
| 987 |
return go.Figure()
|
|
@@ -1020,7 +1050,7 @@ def run_catalog(
|
|
| 1020 |
df = get_dataset(dataset_name, sample_size, seed)
|
| 1021 |
tabfm_params = resolve_tabfm_params(tabfm_preset, tabfm_n_estimators, tabfm_max_rows, tabfm_max_features, tabfm_batch_size, tabfm_enable_nnls, tabfm_crosses, tabfm_svd, tabfm_max_eval_rows)
|
| 1022 |
results, preview, summary = benchmark_frame(df, spec.target, spec.task, sample_size, test_percent / 100, seed, selected_models, include_tabfm, tabfm_params)
|
| 1023 |
-
return summary, results.round(4), metric_chart(results, selected_metrics, chart_style), time_chart(results), preview
|
| 1024 |
|
| 1025 |
|
| 1026 |
def run_upload(
|
|
@@ -1050,7 +1080,7 @@ def run_upload(
|
|
| 1050 |
selected_task = None if task == "Auto" else task.lower()
|
| 1051 |
tabfm_params = resolve_tabfm_params(tabfm_preset, tabfm_n_estimators, tabfm_max_rows, tabfm_max_features, tabfm_batch_size, tabfm_enable_nnls, tabfm_crosses, tabfm_svd, tabfm_max_eval_rows)
|
| 1052 |
results, preview, summary = benchmark_frame(df, target, selected_task, sample_size, test_percent / 100, seed, selected_models, include_tabfm, tabfm_params)
|
| 1053 |
-
return summary, results.round(4), metric_chart(results, selected_metrics, chart_style), time_chart(results), preview
|
| 1054 |
|
| 1055 |
|
| 1056 |
def redraw_metric_chart(results: pd.DataFrame, selected_metrics: list[str], chart_style: str):
|
|
@@ -1059,6 +1089,12 @@ def redraw_metric_chart(results: pd.DataFrame, selected_metrics: list[str], char
|
|
| 1059 |
return metric_chart(pd.DataFrame(results), selected_metrics, chart_style)
|
| 1060 |
|
| 1061 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1062 |
def catalog_table() -> pd.DataFrame:
|
| 1063 |
return pd.DataFrame(
|
| 1064 |
[
|
|
@@ -1136,6 +1172,7 @@ def build_app() -> gr.Blocks:
|
|
| 1136 |
with gr.Row():
|
| 1137 |
chart = gr.Plot(label="Metric comparison")
|
| 1138 |
speed = gr.Plot(label="Speed")
|
|
|
|
| 1139 |
preview = gr.Dataframe(label="Held-out preview", interactive=False)
|
| 1140 |
run_inputs = [
|
| 1141 |
dataset,
|
|
@@ -1156,10 +1193,12 @@ def build_app() -> gr.Blocks:
|
|
| 1156 |
tabfm_svd,
|
| 1157 |
tabfm_max_eval_rows,
|
| 1158 |
]
|
| 1159 |
-
|
|
|
|
| 1160 |
metric_toggles.change(redraw_metric_chart, [leaderboard, metric_toggles, chart_style], chart)
|
|
|
|
| 1161 |
chart_style.change(redraw_metric_chart, [leaderboard, metric_toggles, chart_style], chart)
|
| 1162 |
-
demo.load(run_catalog, run_inputs,
|
| 1163 |
with gr.Tab("Upload Dataset"):
|
| 1164 |
with gr.Row():
|
| 1165 |
with gr.Column(scale=1, elem_classes=["panel"]):
|
|
@@ -1190,6 +1229,7 @@ def build_app() -> gr.Blocks:
|
|
| 1190 |
with gr.Row():
|
| 1191 |
upload_chart = gr.Plot(label="Metric comparison")
|
| 1192 |
upload_speed = gr.Plot(label="Speed")
|
|
|
|
| 1193 |
upload_preview = gr.Dataframe(label="Held-out preview", interactive=False)
|
| 1194 |
upload_btn.click(
|
| 1195 |
run_upload,
|
|
@@ -1214,9 +1254,10 @@ def build_app() -> gr.Blocks:
|
|
| 1214 |
upload_tabfm_svd,
|
| 1215 |
upload_tabfm_max_eval_rows,
|
| 1216 |
],
|
| 1217 |
-
[upload_summary, upload_leaderboard, upload_chart, upload_speed, upload_preview],
|
| 1218 |
)
|
| 1219 |
upload_metric_toggles.change(redraw_metric_chart, [upload_leaderboard, upload_metric_toggles, upload_chart_style], upload_chart)
|
|
|
|
| 1220 |
upload_chart_style.change(redraw_metric_chart, [upload_leaderboard, upload_metric_toggles, upload_chart_style], upload_chart)
|
| 1221 |
with gr.Tab("Dataset Catalog"):
|
| 1222 |
gr.Dataframe(catalog_table(), interactive=False, label="Included benchmark catalog")
|
|
|
|
| 982 |
return fig
|
| 983 |
|
| 984 |
|
| 985 |
+
def bar_chart(results: pd.DataFrame, selected_metrics: list[str] | None = None) -> go.Figure:
|
| 986 |
+
if results is None or results.empty:
|
| 987 |
+
return go.Figure()
|
| 988 |
+
selected_metrics = selected_metrics or METRIC_CHOICES
|
| 989 |
+
metric_cols = [c for c in selected_metrics if c in results.columns and results[c].notna().any()]
|
| 990 |
+
if not metric_cols:
|
| 991 |
+
metric_cols = [c for c in METRIC_CHOICES if c in results.columns and results[c].notna().any()]
|
| 992 |
+
if not metric_cols:
|
| 993 |
+
return go.Figure()
|
| 994 |
+
clean = results.sort_values("rank") if "rank" in results.columns else results.copy()
|
| 995 |
+
long = clean.melt(id_vars=["model"], value_vars=metric_cols, var_name="metric", value_name="score")
|
| 996 |
+
fig = px.bar(
|
| 997 |
+
long,
|
| 998 |
+
x="model",
|
| 999 |
+
y="score",
|
| 1000 |
+
color="metric",
|
| 1001 |
+
barmode="group",
|
| 1002 |
+
color_discrete_sequence=GOOGLE_COLORS,
|
| 1003 |
+
title="Grouped metric comparison",
|
| 1004 |
+
)
|
| 1005 |
+
fig.update_layout(
|
| 1006 |
+
template="plotly_white",
|
| 1007 |
+
height=360,
|
| 1008 |
+
margin=dict(l=25, r=20, t=45, b=35),
|
| 1009 |
+
legend_title_text="Metric",
|
| 1010 |
+
hovermode="x unified",
|
| 1011 |
+
)
|
| 1012 |
+
return fig
|
| 1013 |
+
|
| 1014 |
+
|
| 1015 |
def time_chart(results: pd.DataFrame) -> go.Figure:
|
| 1016 |
if results is None or results.empty or "seconds" not in results:
|
| 1017 |
return go.Figure()
|
|
|
|
| 1050 |
df = get_dataset(dataset_name, sample_size, seed)
|
| 1051 |
tabfm_params = resolve_tabfm_params(tabfm_preset, tabfm_n_estimators, tabfm_max_rows, tabfm_max_features, tabfm_batch_size, tabfm_enable_nnls, tabfm_crosses, tabfm_svd, tabfm_max_eval_rows)
|
| 1052 |
results, preview, summary = benchmark_frame(df, spec.target, spec.task, sample_size, test_percent / 100, seed, selected_models, include_tabfm, tabfm_params)
|
| 1053 |
+
return summary, results.round(4), metric_chart(results, selected_metrics, chart_style), time_chart(results), bar_chart(results, selected_metrics), preview
|
| 1054 |
|
| 1055 |
|
| 1056 |
def run_upload(
|
|
|
|
| 1080 |
selected_task = None if task == "Auto" else task.lower()
|
| 1081 |
tabfm_params = resolve_tabfm_params(tabfm_preset, tabfm_n_estimators, tabfm_max_rows, tabfm_max_features, tabfm_batch_size, tabfm_enable_nnls, tabfm_crosses, tabfm_svd, tabfm_max_eval_rows)
|
| 1082 |
results, preview, summary = benchmark_frame(df, target, selected_task, sample_size, test_percent / 100, seed, selected_models, include_tabfm, tabfm_params)
|
| 1083 |
+
return summary, results.round(4), metric_chart(results, selected_metrics, chart_style), time_chart(results), bar_chart(results, selected_metrics), preview
|
| 1084 |
|
| 1085 |
|
| 1086 |
def redraw_metric_chart(results: pd.DataFrame, selected_metrics: list[str], chart_style: str):
|
|
|
|
| 1089 |
return metric_chart(pd.DataFrame(results), selected_metrics, chart_style)
|
| 1090 |
|
| 1091 |
|
| 1092 |
+
def redraw_bar_chart(results: pd.DataFrame, selected_metrics: list[str]):
|
| 1093 |
+
if results is None or len(results) == 0:
|
| 1094 |
+
return go.Figure()
|
| 1095 |
+
return bar_chart(pd.DataFrame(results), selected_metrics)
|
| 1096 |
+
|
| 1097 |
+
|
| 1098 |
def catalog_table() -> pd.DataFrame:
|
| 1099 |
return pd.DataFrame(
|
| 1100 |
[
|
|
|
|
| 1172 |
with gr.Row():
|
| 1173 |
chart = gr.Plot(label="Metric comparison")
|
| 1174 |
speed = gr.Plot(label="Speed")
|
| 1175 |
+
bars = gr.Plot(label="Grouped comparison")
|
| 1176 |
preview = gr.Dataframe(label="Held-out preview", interactive=False)
|
| 1177 |
run_inputs = [
|
| 1178 |
dataset,
|
|
|
|
| 1193 |
tabfm_svd,
|
| 1194 |
tabfm_max_eval_rows,
|
| 1195 |
]
|
| 1196 |
+
run_outputs = [summary, leaderboard, chart, speed, bars, preview]
|
| 1197 |
+
run_btn.click(run_catalog, run_inputs, run_outputs)
|
| 1198 |
metric_toggles.change(redraw_metric_chart, [leaderboard, metric_toggles, chart_style], chart)
|
| 1199 |
+
metric_toggles.change(redraw_bar_chart, [leaderboard, metric_toggles], bars)
|
| 1200 |
chart_style.change(redraw_metric_chart, [leaderboard, metric_toggles, chart_style], chart)
|
| 1201 |
+
demo.load(run_catalog, run_inputs, run_outputs)
|
| 1202 |
with gr.Tab("Upload Dataset"):
|
| 1203 |
with gr.Row():
|
| 1204 |
with gr.Column(scale=1, elem_classes=["panel"]):
|
|
|
|
| 1229 |
with gr.Row():
|
| 1230 |
upload_chart = gr.Plot(label="Metric comparison")
|
| 1231 |
upload_speed = gr.Plot(label="Speed")
|
| 1232 |
+
upload_bars = gr.Plot(label="Grouped comparison")
|
| 1233 |
upload_preview = gr.Dataframe(label="Held-out preview", interactive=False)
|
| 1234 |
upload_btn.click(
|
| 1235 |
run_upload,
|
|
|
|
| 1254 |
upload_tabfm_svd,
|
| 1255 |
upload_tabfm_max_eval_rows,
|
| 1256 |
],
|
| 1257 |
+
[upload_summary, upload_leaderboard, upload_chart, upload_speed, upload_bars, upload_preview],
|
| 1258 |
)
|
| 1259 |
upload_metric_toggles.change(redraw_metric_chart, [upload_leaderboard, upload_metric_toggles, upload_chart_style], upload_chart)
|
| 1260 |
+
upload_metric_toggles.change(redraw_bar_chart, [upload_leaderboard, upload_metric_toggles], upload_bars)
|
| 1261 |
upload_chart_style.change(redraw_metric_chart, [upload_leaderboard, upload_metric_toggles, upload_chart_style], upload_chart)
|
| 1262 |
with gr.Tab("Dataset Catalog"):
|
| 1263 |
gr.Dataframe(catalog_table(), interactive=False, label="Included benchmark catalog")
|