| 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()) |
|
|