misscp / src /sepsis_mcp /appendix_phase2_analysis.py
Anonymous
Initial anonymous MissCP release
32f5a65
Raw
History Blame Contribute Delete
11.9 kB
from __future__ import annotations
import argparse
import json
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
import pandas as pd
DEFAULT_METHODS = ("standard", "missingness_aware", "weighted_missingness_aware")
@dataclass
class AppendixPhase2Config:
alpha_sweep_dirs: tuple[Path, ...]
grouping_sweep_dir: Path
selection_sweep_dir: Path
output_dir: Path
methods: tuple[str, ...] = DEFAULT_METHODS
model_type: str = "xgboost"
def _load_csv(path: Path) -> pd.DataFrame:
return pd.read_csv(path) if path.exists() else pd.DataFrame()
def build_alpha_sweep_summary(
*,
repeated_summary_frames: list[pd.DataFrame],
methods: tuple[str, ...],
model_type: str,
) -> pd.DataFrame:
if not repeated_summary_frames:
return pd.DataFrame()
combined = pd.concat(repeated_summary_frames, ignore_index=True)
filtered = combined[
combined["method"].isin(methods)
& combined["model_type"].eq(model_type)
].copy()
if filtered.empty:
return filtered
preferred_columns = [
"alpha",
"method",
"run_count",
"empirical_coverage_mean",
"empirical_coverage_std",
"max_group_coverage_gap_mean",
"max_group_coverage_gap_std",
"average_set_size_mean",
"average_set_size_std",
"worst_hospital_coverage_mean",
"worst_hospital_coverage_std",
"wcr",
"source_dir",
]
filtered = filtered.sort_values(["alpha", "method"]).reset_index(drop=True)
return filtered[[column for column in preferred_columns if column in filtered.columns]]
def build_grouping_baselines_summary(
*,
repeated_summary: pd.DataFrame,
methods: tuple[str, ...],
model_type: str,
) -> pd.DataFrame:
filtered = repeated_summary[
repeated_summary["method"].isin(methods)
& repeated_summary["model_type"].eq(model_type)
].copy()
if filtered.empty:
return filtered
base_columns = [
"missingness_grouping_strategy",
"model_type",
"alpha",
"selection_fraction",
"mask_strategy",
"mask_rate",
"weighted_shrinkage_lambda",
]
standard = (
filtered[filtered["method"] == "standard"][base_columns + ["max_group_coverage_gap_mean"]]
.rename(columns={"max_group_coverage_gap_mean": "standard_gap_mean"})
)
summary = filtered.merge(standard, on=base_columns, how="left", validate="many_to_one")
summary["gap_reduction_vs_standard_mean"] = (
summary["standard_gap_mean"] - summary["max_group_coverage_gap_mean"]
)
preferred_columns = [
"missingness_grouping_strategy",
"method",
"alpha",
"run_count",
"empirical_coverage_mean",
"max_group_coverage_gap_mean",
"gap_reduction_vs_standard_mean",
"average_set_size_mean",
"worst_hospital_coverage_mean",
"wcr",
]
return summary.sort_values(
["missingness_grouping_strategy", "method"]
).reset_index(drop=True)[preferred_columns]
def build_selection_stability_summary(
*,
selected_variable_stability: pd.DataFrame,
selected_variable_jaccard: pd.DataFrame,
repeated_summary: pd.DataFrame,
model_type: str,
) -> pd.DataFrame:
if selected_variable_stability.empty:
return pd.DataFrame()
filtered = selected_variable_stability[selected_variable_stability["model_type"].eq(model_type)].copy()
if filtered.empty:
return filtered
base_columns = [
"experiment",
"method",
"min_hospital_admissions",
"alpha",
"selection_fraction",
"model_type",
"missingness_grouping_strategy",
"mask_strategy",
"mask_rate",
"selective_feature_group",
"weighted_shrinkage_lambda",
]
grouped = filtered.groupby(base_columns, dropna=False)
rows: list[dict[str, Any]] = []
for key, frame in grouped:
top = frame.sort_values(["frequency", "count", "selected_variable"], ascending=[False, False, True]).iloc[0]
top3_mass = float(frame.sort_values(["frequency", "count"], ascending=[False, False]).head(3)["frequency"].sum())
row = dict(zip(base_columns, key if isinstance(key, tuple) else (key,), strict=False))
row.update(
{
"top_selected_variable": str(top["selected_variable"]),
"top_selected_variable_frequency": float(top["frequency"]),
"top3_frequency_mass": top3_mass,
"unique_selected_variables": int(frame["selected_variable"].nunique()),
}
)
rows.append(row)
summary = pd.DataFrame(rows)
if not selected_variable_jaccard.empty:
summary = summary.merge(
selected_variable_jaccard[
base_columns + ["run_count", "mean_pairwise_jaccard", "min_pairwise_jaccard"]
],
on=base_columns,
how="left",
validate="one_to_one",
)
if not repeated_summary.empty:
repeated_filtered = repeated_summary[repeated_summary["model_type"].eq(model_type)].copy()
keep = base_columns + [
"empirical_coverage_mean",
"max_group_coverage_gap_mean",
"average_set_size_mean",
"wcr",
]
summary = summary.merge(
repeated_filtered[keep],
on=base_columns,
how="left",
validate="one_to_one",
)
return summary.sort_values(["method", "alpha"]).reset_index(drop=True)
def build_selection_performance_by_variable(
*,
overall_summary: pd.DataFrame,
model_type: str,
) -> pd.DataFrame:
if overall_summary.empty or "selected_variable" not in overall_summary.columns:
return pd.DataFrame()
filtered = overall_summary[
overall_summary["model_type"].eq(model_type)
& overall_summary["selected_variable"].notna()
].copy()
if filtered.empty:
return filtered
group_columns = ["method", "selected_variable"]
aggregated = filtered.groupby(group_columns, dropna=False, as_index=False).agg(
run_count=("run_id", "nunique"),
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"),
worst_hospital_coverage_mean=("worst_hospital_coverage", "mean"),
worst_hospital_coverage_std=("worst_hospital_coverage", "std"),
)
aggregated["empirical_coverage_std"] = aggregated["empirical_coverage_std"].fillna(0.0)
aggregated["max_group_coverage_gap_std"] = aggregated["max_group_coverage_gap_std"].fillna(0.0)
aggregated["average_set_size_std"] = aggregated["average_set_size_std"].fillna(0.0)
aggregated["worst_hospital_coverage_std"] = aggregated["worst_hospital_coverage_std"].fillna(0.0)
return aggregated.sort_values(["method", "max_group_coverage_gap_mean", "selected_variable"]).reset_index(drop=True)
def run_appendix_phase2_analysis(config: AppendixPhase2Config) -> dict[str, Path]:
config.output_dir.mkdir(parents=True, exist_ok=True)
alpha_repeated_frames: list[pd.DataFrame] = []
for sweep_dir in config.alpha_sweep_dirs:
frame = _load_csv(sweep_dir / "repeated_summary.csv")
if frame.empty:
continue
frame = frame.copy()
frame["source_dir"] = str(sweep_dir)
alpha_repeated_frames.append(frame)
grouping_repeated_summary = _load_csv(config.grouping_sweep_dir / "repeated_summary.csv")
selection_stability = _load_csv(config.selection_sweep_dir / "selected_variable_stability.csv")
selection_jaccard = _load_csv(config.selection_sweep_dir / "selected_variable_jaccard_summary.csv")
selection_repeated = _load_csv(config.selection_sweep_dir / "repeated_summary.csv")
selection_overall = _load_csv(config.selection_sweep_dir / "overall_summary.csv")
alpha_summary = build_alpha_sweep_summary(
repeated_summary_frames=alpha_repeated_frames,
methods=config.methods,
model_type=config.model_type,
)
grouping_summary = build_grouping_baselines_summary(
repeated_summary=grouping_repeated_summary,
methods=config.methods,
model_type=config.model_type,
)
stability_summary = build_selection_stability_summary(
selected_variable_stability=selection_stability,
selected_variable_jaccard=selection_jaccard,
repeated_summary=selection_repeated,
model_type=config.model_type,
)
performance_by_variable = build_selection_performance_by_variable(
overall_summary=selection_overall,
model_type=config.model_type,
)
alpha_path = config.output_dir / "alpha_sweep_summary.csv"
grouping_path = config.output_dir / "grouping_baselines_summary.csv"
stability_path = config.output_dir / "selection_stability_summary.csv"
performance_path = config.output_dir / "selection_performance_by_variable.csv"
config_path = config.output_dir / "config.json"
manifest_path = config.output_dir / "manifest.json"
alpha_summary.to_csv(alpha_path, index=False)
grouping_summary.to_csv(grouping_path, index=False)
stability_summary.to_csv(stability_path, index=False)
performance_by_variable.to_csv(performance_path, index=False)
config_path.write_text(json.dumps(asdict(config), indent=2, default=str), encoding="utf-8")
manifest_path.write_text(
json.dumps(
{
"alpha_sweep_summary": str(alpha_path),
"grouping_baselines_summary": str(grouping_path),
"selection_stability_summary": str(stability_path),
"selection_performance_by_variable": str(performance_path),
"config": str(config_path),
},
indent=2,
sort_keys=True,
),
encoding="utf-8",
)
return {
"alpha_sweep_summary": alpha_path,
"grouping_baselines_summary": grouping_path,
"selection_stability_summary": stability_path,
"selection_performance_by_variable": performance_path,
"config": config_path,
"manifest": manifest_path,
}
def _parse_path_list(value: str) -> tuple[Path, ...]:
return tuple(Path(item.strip()) for item in value.split(",") if item.strip())
def _parse_methods(value: str) -> tuple[str, ...]:
return tuple(item.strip() for item in value.split(",") if item.strip())
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(prog="sepsis-mcp-appendix-phase2-analysis")
parser.add_argument("--alpha-sweep-dirs", type=_parse_path_list, required=True)
parser.add_argument("--grouping-sweep-dir", type=Path, required=True)
parser.add_argument("--selection-sweep-dir", type=Path, required=True)
parser.add_argument("--methods", type=_parse_methods, default=",".join(DEFAULT_METHODS))
parser.add_argument("--model-type", default="xgboost")
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_appendix_phase2_analysis(
AppendixPhase2Config(
alpha_sweep_dirs=args.alpha_sweep_dirs,
grouping_sweep_dir=args.grouping_sweep_dir,
selection_sweep_dir=args.selection_sweep_dir,
output_dir=args.output_dir,
methods=args.methods,
model_type=args.model_type,
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())