Skip to content

Commit bdaeb70

Browse files
committed
fix: aggregate repeated-split stability evidence
1 parent e0e0fd0 commit bdaeb70

1 file changed

Lines changed: 102 additions & 40 deletions

File tree

app/pages/3_Model_Benchmark_Selection.py

Lines changed: 102 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -441,64 +441,126 @@ def build_threshold_chart(
441441
def build_stability_chart(
442442
stability: pd.DataFrame,
443443
) -> object | None:
444-
if stability.empty or "model_name" not in stability.columns:
445-
return None
444+
"""Build repeated-split stability from raw seed rows or summary rows.
446445
447-
mean_column = None
448-
std_column = None
446+
The saved evidence currently contains one row per model and random seed:
447+
model_name, seed, pr_auc, and roc_auc. The chart therefore calculates the
448+
model-level mean and standard deviation before drawing the bars.
449449
450-
for candidate in (
451-
"mean_pr_auc",
452-
"pr_auc_mean",
453-
"average_pr_auc",
454-
):
455-
if candidate in stability.columns:
456-
mean_column = candidate
457-
break
450+
If a future pipeline saves pre-aggregated mean/std columns, those columns
451+
are also supported.
452+
"""
458453

459-
for candidate in (
460-
"std_pr_auc",
461-
"pr_auc_std",
462-
"stdev_pr_auc",
463-
):
464-
if candidate in stability.columns:
465-
std_column = candidate
466-
break
467-
468-
if mean_column is None:
454+
if stability.empty or "model_name" not in stability.columns:
469455
return None
470456

471-
chart_data = stability.copy()
472-
chart_data["Mean PR-AUC"] = pd.to_numeric(
473-
chart_data[mean_column],
474-
errors="coerce",
457+
data = stability.copy()
458+
459+
summary_mean_column = next(
460+
(
461+
column
462+
for column in (
463+
"mean_pr_auc",
464+
"pr_auc_mean",
465+
"average_pr_auc",
466+
)
467+
if column in data.columns
468+
),
469+
None,
475470
)
476471

477-
error_column = None
472+
summary_std_column = next(
473+
(
474+
column
475+
for column in (
476+
"std_pr_auc",
477+
"pr_auc_std",
478+
"stdev_pr_auc",
479+
)
480+
if column in data.columns
481+
),
482+
None,
483+
)
484+
485+
if summary_mean_column is not None:
486+
summary = pd.DataFrame(
487+
{
488+
"model_name": data["model_name"].astype(str),
489+
"mean_pr_auc": pd.to_numeric(
490+
data[summary_mean_column],
491+
errors="coerce",
492+
),
493+
}
494+
)
478495

479-
if std_column is not None:
480-
chart_data["PR-AUC standard deviation"] = pd.to_numeric(
481-
chart_data[std_column],
496+
if summary_std_column is not None:
497+
summary["std_pr_auc"] = pd.to_numeric(
498+
data[summary_std_column],
499+
errors="coerce",
500+
)
501+
else:
502+
summary["std_pr_auc"] = 0.0
503+
504+
elif "pr_auc" in data.columns:
505+
# Raw repeated-split evidence: aggregate one row per model.
506+
data["pr_auc"] = pd.to_numeric(
507+
data["pr_auc"],
482508
errors="coerce",
483509
)
484-
error_column = "PR-AUC standard deviation"
510+
data = data.dropna(
511+
subset=["model_name", "pr_auc"],
512+
)
485513

486-
chart_data = chart_data.dropna(subset=["Mean PR-AUC"])
487-
chart_data = chart_data.sort_values(
488-
"Mean PR-AUC",
489-
ascending=True,
514+
if data.empty:
515+
return None
516+
517+
summary = (
518+
data.groupby(
519+
"model_name",
520+
as_index=False,
521+
dropna=False,
522+
)
523+
.agg(
524+
mean_pr_auc=("pr_auc", "mean"),
525+
std_pr_auc=("pr_auc", "std"),
526+
repeated_splits=("pr_auc", "count"),
527+
)
528+
)
529+
summary["std_pr_auc"] = (
530+
summary["std_pr_auc"].fillna(0.0)
531+
)
532+
533+
else:
534+
return None
535+
536+
summary = summary.dropna(
537+
subset=["model_name", "mean_pr_auc"],
490538
)
491539

492-
if chart_data.empty:
540+
if summary.empty:
493541
return None
494542

543+
summary = summary.sort_values(
544+
"mean_pr_auc",
545+
ascending=True,
546+
)
547+
548+
hover_data = {
549+
"mean_pr_auc": ":.3f",
550+
"std_pr_auc": ":.3f",
551+
}
552+
553+
if "repeated_splits" in summary.columns:
554+
hover_data["repeated_splits"] = True
555+
495556
figure = px.bar(
496-
chart_data,
497-
x="Mean PR-AUC",
557+
summary,
558+
x="mean_pr_auc",
498559
y="model_name",
499560
orientation="h",
500-
error_x=error_column,
501-
title="Repeated-split stability evidence",
561+
error_x="std_pr_auc",
562+
hover_data=hover_data,
563+
title="Repeated-split stability: mean PR-AUC by model",
502564
)
503565
figure.update_layout(
504566
height=430,

0 commit comments

Comments
 (0)