misscp / src /sepsis_mcp /appendix_multigroup_ablation.py
Anonymous
Initial anonymous MissCP release
32f5a65
Raw
History Blame Contribute Delete
7.99 kB
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())