| 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__": |
| raise SystemExit(main()) |
|
|