from __future__ import annotations import argparse import json from dataclasses import asdict, dataclass from pathlib import Path import pandas as pd from sepsis_mcp.gossis_experiment import GossisRunConfig, run_gossis_experiment MULTIGROUP_VARIANTS = ( ("binary_selected", {"missingness_grouping_strategy": "coverage_gap_variable"}), ("mask_cluster_k3", {"missingness_grouping_strategy": "mask_cluster", "mask_cluster_k_grid": [3]}), ("mask_cluster_k4", {"missingness_grouping_strategy": "mask_cluster", "mask_cluster_k_grid": [4]}), ("top2_cartesian", {"missingness_grouping_strategy": "top_missingness_cartesian", "composite_top_l": 2}), ) @dataclass class MultiGroupAblationConfig: data_root: Path output_dir: Path random_state_grid: tuple[int, ...] = (0, 1, 2) model_type: str = "xgboost" alpha: float = 0.1 selection_fraction: float = 0.1 min_hospital_admissions: int = 500 min_selection_group_rows: int = 100 methods: tuple[str, ...] = ("standard", "missingness_aware") def build_multigroup_summary( *, overall_summary: pd.DataFrame, subgroup_summary: pd.DataFrame, ) -> pd.DataFrame: if overall_summary.empty: return pd.DataFrame() subgroup_minima = ( subgroup_summary.groupby(["run_id", "grouping_variant", "method"], as_index=False) .agg(smallest_group_size=("count", "min")) ) merged = overall_summary.merge( subgroup_minima, on=["run_id", "grouping_variant", "method"], how="left", validate="one_to_one", ) summary = ( merged.groupby(["grouping_variant", "method"], as_index=False) .agg( run_count=("run_id", "nunique"), group_count_mean=("group_count", "mean"), smallest_group_size_mean=("smallest_group_size", "mean"), smallest_group_size_min=("smallest_group_size", "min"), empirical_coverage_mean=("empirical_coverage", "mean"), empirical_coverage_std=("empirical_coverage", "std"), max_group_coverage_gap_mean=("max_group_coverage_gap", "mean"), max_group_coverage_gap_std=("max_group_coverage_gap", "std"), average_set_size_mean=("average_set_size", "mean"), average_set_size_std=("average_set_size", "std"), ) .sort_values(["method", "group_count_mean", "grouping_variant"]) .reset_index(drop=True) ) for column in summary.columns: if column.endswith("_std"): summary[column] = summary[column].fillna(0.0) return summary def write_multigroup_outputs(*, summary: pd.DataFrame, output_dir: Path) -> dict[str, Path]: output_dir.mkdir(parents=True, exist_ok=True) summary_path = output_dir / "multigroup_ablation_summary.csv" manifest_path = output_dir / "manifest.json" summary.to_csv(summary_path, index=False) manifest_path.write_text( json.dumps({"summary": str(summary_path)}, indent=2, sort_keys=True), encoding="utf-8", ) return {"summary": summary_path, "manifest": manifest_path} def run_multigroup_ablation(config: MultiGroupAblationConfig) -> dict[str, Path]: config.output_dir.mkdir(parents=True, exist_ok=True) overall_rows: list[pd.DataFrame] = [] subgroup_rows: list[pd.DataFrame] = [] for random_state in config.random_state_grid: for grouping_variant, variant_kwargs in MULTIGROUP_VARIANTS: run_dir = config.output_dir / "runs" / f"{grouping_variant}__seed{random_state}" overall_path = run_dir / "overall_summary.csv" subgroup_path = run_dir / "subgroup_coverage.csv" if overall_path.exists() and subgroup_path.exists(): outputs = { "overall_summary_path": overall_path, "subgroup_path": subgroup_path, } else: outputs = run_gossis_experiment( GossisRunConfig( data_root=config.data_root, output_dir=run_dir, random_state=int(random_state), model_type=config.model_type, alpha=config.alpha, selection_fraction=config.selection_fraction, min_hospital_admissions=config.min_hospital_admissions, min_selection_group_rows=config.min_selection_group_rows, **variant_kwargs, ) ) overall = pd.read_csv(outputs["overall_summary_path"]) subgroup = pd.read_csv(outputs["subgroup_path"]) run_id = run_dir.name overall["run_id"] = run_id subgroup["run_id"] = run_id overall["grouping_variant"] = grouping_variant subgroup["grouping_variant"] = grouping_variant overall_rows.append(overall[overall["method"].isin(config.methods)].copy()) subgroup_rows.append(subgroup[subgroup["method"].isin(config.methods)].copy()) overall_combined = pd.concat(overall_rows, ignore_index=True) subgroup_combined = pd.concat(subgroup_rows, ignore_index=True) summary = build_multigroup_summary( overall_summary=overall_combined, subgroup_summary=subgroup_combined, ) overall_path = config.output_dir / "overall_summary.csv" subgroup_path = config.output_dir / "subgroup_summary.csv" config_path = config.output_dir / "config.json" manifest_path = config.output_dir / "manifest.json" summary_path = config.output_dir / "multigroup_ablation_summary.csv" overall_combined.to_csv(overall_path, index=False) subgroup_combined.to_csv(subgroup_path, index=False) summary.to_csv(summary_path, index=False) config_path.write_text(json.dumps(asdict(config), indent=2, default=str), encoding="utf-8") manifest_path.write_text( json.dumps( { "overall_summary": str(overall_path), "subgroup_summary": str(subgroup_path), "summary": str(summary_path), "config": str(config_path), }, indent=2, sort_keys=True, ), encoding="utf-8", ) return { "overall_summary": overall_path, "subgroup_summary": subgroup_path, "summary": summary_path, "config": config_path, "manifest": manifest_path, } def _parse_int_grid(value: str) -> tuple[int, ...]: return tuple(int(item.strip()) for item in value.split(",") if item.strip()) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="sepsis-mcp-appendix-multigroup-ablation") parser.add_argument("--data-root", type=Path, required=True) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--random-state-grid", type=_parse_int_grid, default="0,1,2") parser.add_argument("--model-type", default="xgboost") parser.add_argument("--alpha", type=float, default=0.1) parser.add_argument("--selection-fraction", type=float, default=0.1) parser.add_argument("--min-hospital-admissions", type=int, default=500) parser.add_argument("--min-selection-group-rows", type=int, default=100) return parser def main(argv: list[str] | None = None) -> int: parser = build_parser() args = parser.parse_args(argv) run_multigroup_ablation( MultiGroupAblationConfig( data_root=args.data_root, output_dir=args.output_dir, random_state_grid=args.random_state_grid, model_type=args.model_type, alpha=args.alpha, selection_fraction=args.selection_fraction, min_hospital_admissions=args.min_hospital_admissions, min_selection_group_rows=args.min_selection_group_rows, ) ) return 0 if __name__ == "__main__": raise SystemExit(main())