misscp / src /sepsis_mcp /external_baselines_analysis.py
Anonymous
Initial anonymous MissCP release
32f5a65
Raw
History Blame Contribute Delete
13.4 kB
from __future__ import annotations
import argparse
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import numpy as np
import pandas as pd
from sklearn.metrics import adjusted_rand_score, normalized_mutual_info_score
OVERLAP_GROUP_COLUMNS = ["split_id", "split", "test_unit", "rotation"]
INTERPRETABILITY_GROUP_COLUMNS = ["split_id", "test_unit", "rotation"]
@dataclass(frozen=True)
class ExternalBaselinesAnalysisConfig:
run_dirs: tuple[Path, ...]
output_dir: Path
def _group_columns(frame: pd.DataFrame, candidates: list[str]) -> list[str]:
return [column for column in candidates if column in frame.columns]
def _run_context(run_dir: Path) -> dict[str, str]:
return {
"run_dir": str(run_dir),
"run_name": run_dir.name,
}
def _load_optional_frame(path: Path) -> pd.DataFrame:
if not path.exists():
return pd.DataFrame()
return pd.read_csv(path)
def _compute_overlap_outputs(run_dir: Path, assignments: pd.DataFrame) -> tuple[pd.DataFrame, pd.DataFrame]:
if assignments.empty:
return pd.DataFrame(), pd.DataFrame()
summary_rows: list[dict[str, Any]] = []
detail_rows: list[dict[str, Any]] = []
group_columns = _group_columns(assignments, OVERLAP_GROUP_COLUMNS)
grouped = assignments.groupby(group_columns, dropna=False) if group_columns else [((), assignments)]
run_context = _run_context(run_dir)
for key, frame in grouped:
context = dict(zip(group_columns, key if isinstance(key, tuple) else (key,), strict=False))
contingency = pd.crosstab(frame["learned_group"], frame["missingness_group"])
learned_totals = contingency.sum(axis=1)
missingness_totals = contingency.sum(axis=0)
learned_purity = (contingency.max(axis=1) / learned_totals).replace([np.inf, -np.inf], np.nan).fillna(0.0)
missingness_purity = (contingency.max(axis=0) / missingness_totals).replace([np.inf, -np.inf], np.nan).fillna(0.0)
weighted_learned_purity = float(np.average(learned_purity.to_numpy(dtype=float), weights=learned_totals.to_numpy(dtype=float)))
weighted_missingness_purity = float(
np.average(missingness_purity.to_numpy(dtype=float), weights=missingness_totals.to_numpy(dtype=float))
)
summary_rows.append(
{
**run_context,
**context,
"n_samples": int(len(frame)),
"learned_group_count": int(frame["learned_group"].nunique()),
"missingness_group_count": int(frame["missingness_group"].nunique()),
"ari": float(
adjusted_rand_score(
frame["learned_group"].to_numpy(dtype=int),
frame["missingness_group"].to_numpy(dtype=int),
)
),
"nmi": float(
normalized_mutual_info_score(
frame["learned_group"].to_numpy(dtype=int),
frame["missingness_group"].to_numpy(dtype=int),
)
),
"weighted_learned_group_purity": weighted_learned_purity,
"weighted_missingness_group_purity": weighted_missingness_purity,
}
)
contingency_long = contingency.stack().rename("overlap_count").reset_index()
for row in contingency_long.itertuples(index=False):
learned_total = int(learned_totals.loc[int(row.learned_group)])
missingness_total = int(missingness_totals.loc[int(row.missingness_group)])
detail_rows.append(
{
**run_context,
**context,
"learned_group": int(row.learned_group),
"missingness_group": int(row.missingness_group),
"overlap_count": int(row.overlap_count),
"learned_group_fraction": float(row.overlap_count / learned_total),
"missingness_group_fraction": float(row.overlap_count / missingness_total),
}
)
return pd.DataFrame(summary_rows), pd.DataFrame(detail_rows)
def _compute_interpretability_output(run_dir: Path, feature_importance: pd.DataFrame) -> pd.DataFrame:
if feature_importance.empty:
return pd.DataFrame()
summary_rows: list[dict[str, Any]] = []
group_columns = _group_columns(feature_importance, INTERPRETABILITY_GROUP_COLUMNS)
grouped = feature_importance.groupby(group_columns, dropna=False) if group_columns else [((), feature_importance)]
run_context = _run_context(run_dir)
for key, frame in grouped:
context = dict(zip(group_columns, key if isinstance(key, tuple) else (key,), strict=False))
sorted_frame = frame.sort_values(["importance", "split_count", "feature"], ascending=[False, False, True]).reset_index(drop=True)
total_importance = float(sorted_frame["importance"].sum())
missingness_frame = sorted_frame[sorted_frame["is_missingness_proxy"].astype(bool)]
top_feature = sorted_frame.iloc[0] if not sorted_frame.empty else None
top_missingness = missingness_frame.iloc[0] if not missingness_frame.empty else None
summary_rows.append(
{
**run_context,
**context,
"used_feature_count": int((sorted_frame["split_count"] > 0).sum()),
"total_importance": total_importance,
"missingness_proxy_importance": float(missingness_frame["importance"].sum()),
"share_missingness_proxy_importance": (
float(missingness_frame["importance"].sum() / total_importance) if total_importance > 0.0 else 0.0
),
"top_feature": None if top_feature is None else str(top_feature["feature"]),
"top_feature_importance": None if top_feature is None else float(top_feature["importance"]),
"top_missingness_proxy_feature": None if top_missingness is None else str(top_missingness["feature"]),
"top_missingness_proxy_importance": None if top_missingness is None else float(top_missingness["importance"]),
}
)
return pd.DataFrame(summary_rows)
def _compute_stress_outputs(run_dir: Path, stress_aggregate: pd.DataFrame) -> tuple[pd.DataFrame, pd.DataFrame]:
if stress_aggregate.empty:
return pd.DataFrame(), pd.DataFrame()
per_drop_rows: list[dict[str, Any]] = []
summary_rows: list[dict[str, Any]] = []
run_context = _run_context(run_dir)
def _round(value: float) -> float:
return float(np.round(value, 12))
for method, frame in stress_aggregate.groupby("method", dropna=False):
sorted_frame = frame.sort_values("drop_rate").reset_index(drop=True)
baseline = sorted_frame.iloc[0]
for row in sorted_frame.itertuples(index=False):
per_drop_rows.append(
{
**run_context,
"method": str(method),
"drop_rate": float(row.drop_rate),
"baseline_drop_rate": float(baseline["drop_rate"]),
"coverage_change": _round(float(row.mean_coverage - baseline["mean_coverage"])),
"gap_change": _round(float(row.mean_gap - baseline["mean_gap"])),
"set_size_change": _round(float(row.mean_set_size - baseline["mean_set_size"])),
}
)
max_drop_row = sorted_frame.iloc[-1]
coverage_changes = sorted_frame["mean_coverage"] - float(baseline["mean_coverage"])
gap_changes = sorted_frame["mean_gap"] - float(baseline["mean_gap"])
set_size_changes = sorted_frame["mean_set_size"] - float(baseline["mean_set_size"])
summary_rows.append(
{
**run_context,
"method": str(method),
"baseline_drop_rate": float(baseline["drop_rate"]),
"baseline_coverage": float(baseline["mean_coverage"]),
"baseline_gap": float(baseline["mean_gap"]),
"baseline_set_size": float(baseline["mean_set_size"]),
"max_drop_rate": float(max_drop_row["drop_rate"]),
"coverage_change_at_max_drop": _round(float(max_drop_row["mean_coverage"] - baseline["mean_coverage"])),
"gap_change_at_max_drop": _round(float(max_drop_row["mean_gap"] - baseline["mean_gap"])),
"set_size_change_at_max_drop": _round(float(max_drop_row["mean_set_size"] - baseline["mean_set_size"])),
"worst_coverage_change": _round(float(coverage_changes.min())),
"worst_gap_change": _round(float(gap_changes.max())),
"largest_set_size_change": _round(float(np.max(np.abs(set_size_changes.to_numpy(dtype=float))))),
}
)
return pd.DataFrame(summary_rows), pd.DataFrame(per_drop_rows)
def run_external_baselines_analysis(
*,
run_dirs: tuple[Path, ...],
output_dir: Path,
) -> dict[str, Path]:
output_dir.mkdir(parents=True, exist_ok=True)
overlap_summaries: list[pd.DataFrame] = []
overlap_details: list[pd.DataFrame] = []
interpretability_summaries: list[pd.DataFrame] = []
stress_summaries: list[pd.DataFrame] = []
stress_by_drop: list[pd.DataFrame] = []
for run_dir in run_dirs:
overlap_summary, overlap_detail = _compute_overlap_outputs(
run_dir,
_load_optional_frame(run_dir / "partition_overlap_assignments.csv"),
)
interpretability = _compute_interpretability_output(
run_dir,
_load_optional_frame(run_dir / "partition_feature_importance.csv"),
)
stress_summary, stress_drop = _compute_stress_outputs(
run_dir,
_load_optional_frame(run_dir / "stress_aggregate.csv"),
)
if not overlap_summary.empty:
overlap_summaries.append(overlap_summary)
if not overlap_detail.empty:
overlap_details.append(overlap_detail)
if not interpretability.empty:
interpretability_summaries.append(interpretability)
if not stress_summary.empty:
stress_summaries.append(stress_summary)
if not stress_drop.empty:
stress_by_drop.append(stress_drop)
overlap_summary_frame = pd.concat(overlap_summaries, ignore_index=True) if overlap_summaries else pd.DataFrame()
overlap_detail_frame = pd.concat(overlap_details, ignore_index=True) if overlap_details else pd.DataFrame()
interpretability_frame = (
pd.concat(interpretability_summaries, ignore_index=True) if interpretability_summaries else pd.DataFrame()
)
stress_summary_frame = pd.concat(stress_summaries, ignore_index=True) if stress_summaries else pd.DataFrame()
stress_by_drop_frame = pd.concat(stress_by_drop, ignore_index=True) if stress_by_drop else pd.DataFrame()
overlap_summary_path = output_dir / "overlap_summary.csv"
overlap_detail_path = output_dir / "group_overlap_details.csv"
interpretability_path = output_dir / "partition_interpretability.csv"
stress_summary_path = output_dir / "stress_degradation_summary.csv"
stress_by_drop_path = output_dir / "stress_degradation_by_drop_rate.csv"
analysis_summary_path = output_dir / "analysis_summary.json"
overlap_summary_frame.to_csv(overlap_summary_path, index=False)
overlap_detail_frame.to_csv(overlap_detail_path, index=False)
interpretability_frame.to_csv(interpretability_path, index=False)
stress_summary_frame.to_csv(stress_summary_path, index=False)
stress_by_drop_frame.to_csv(stress_by_drop_path, index=False)
analysis_summary = {
"run_count": len(run_dirs),
"runs_with_overlap": int(overlap_summary_frame["run_name"].nunique()) if not overlap_summary_frame.empty else 0,
"runs_with_stress": int(stress_summary_frame["run_name"].nunique()) if not stress_summary_frame.empty else 0,
"mean_ari": None if overlap_summary_frame.empty else float(overlap_summary_frame["ari"].mean()),
"mean_nmi": None if overlap_summary_frame.empty else float(overlap_summary_frame["nmi"].mean()),
}
analysis_summary_path.write_text(json.dumps(analysis_summary, indent=2, sort_keys=True), encoding="utf-8")
return {
"overlap_summary": overlap_summary_path,
"overlap_details": overlap_detail_path,
"partition_interpretability": interpretability_path,
"stress_degradation_summary": stress_summary_path,
"stress_degradation_by_drop_rate": stress_by_drop_path,
"analysis_summary": analysis_summary_path,
}
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="Analyze learned-partition overlap and stress degradation outputs")
parser.add_argument("--run-dirs", type=Path, nargs="+", required=True)
parser.add_argument("--output-dir", type=Path, required=True)
return parser
def main(argv: list[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
run_external_baselines_analysis(
run_dirs=tuple(args.run_dirs),
output_dir=args.output_dir,
)
return 0
if __name__ == "__main__": # pragma: no cover
raise SystemExit(main())