misscp / src /sepsis_mcp /appendix_weighted_baselines_analysis.py
Anonymous
Initial anonymous MissCP release
32f5a65
Raw
History Blame Contribute Delete
5.87 kB
from __future__ import annotations
import argparse
import json
from dataclasses import asdict, dataclass
from pathlib import Path
import pandas as pd
WEIGHTED_BASELINE_METHODS = (
"standard",
"missingness_aware",
"weighted_missingness_aware",
"shrunk_weighted_missingness_aware",
"weighted_propensity_logistic",
"weighted_propensity_tree",
)
@dataclass
class WeightedBaselinesAnalysisConfig:
sweep_dir: Path
output_dir: Path
model_type: str = "xgboost"
methods: tuple[str, ...] = WEIGHTED_BASELINE_METHODS
def run_weighted_baselines_analysis(config: WeightedBaselinesAnalysisConfig) -> dict[str, Path]:
config.output_dir.mkdir(parents=True, exist_ok=True)
repeated_summary = pd.read_csv(config.sweep_dir / "repeated_summary.csv")
subgroup_summary = pd.read_csv(config.sweep_dir / "subgroup_summary.csv")
repeated_filtered = repeated_summary[
repeated_summary["method"].isin(config.methods)
& repeated_summary["model_type"].eq(config.model_type)
].copy()
subgroup_filtered = subgroup_summary[
subgroup_summary["method"].isin(config.methods)
& subgroup_summary["model_type"].eq(config.model_type)
].copy()
if not repeated_filtered.empty:
standard = (
repeated_filtered[repeated_filtered["method"] == "standard"][
[
"experiment",
"min_hospital_admissions",
"alpha",
"selection_fraction",
"model_type",
"missingness_grouping_strategy",
"mask_strategy",
"mask_rate",
"selective_feature_group",
"weighted_shrinkage_lambda",
"max_group_coverage_gap_mean",
]
]
.rename(columns={"max_group_coverage_gap_mean": "standard_gap_mean"})
)
join_columns = [column for column in standard.columns if column != "standard_gap_mean"]
repeated_filtered = repeated_filtered.merge(
standard,
on=join_columns,
how="left",
validate="many_to_one",
)
repeated_filtered["gap_reduction_vs_standard_mean"] = (
repeated_filtered["standard_gap_mean"] - repeated_filtered["max_group_coverage_gap_mean"]
)
repeated_filtered = repeated_filtered.sort_values(
["max_group_coverage_gap_mean", "average_set_size_mean", "method"]
).reset_index(drop=True)
subgroup_aggregate = pd.DataFrame()
if not subgroup_filtered.empty:
aggregations: dict[str, tuple[str, str]] = {
"run_count": ("run_id", "nunique"),
"count_mean": ("count", "mean"),
"coverage_mean": ("coverage", "mean"),
"coverage_std": ("coverage", "std"),
"coverage_gap_mean": ("coverage_gap", "mean"),
"coverage_gap_std": ("coverage_gap", "std"),
}
if "average_set_size" in subgroup_filtered.columns:
aggregations["average_set_size_mean"] = ("average_set_size", "mean")
aggregations["average_set_size_std"] = ("average_set_size", "std")
subgroup_aggregate = (
subgroup_filtered.groupby(["method", "group_label"], as_index=False)
.agg(**aggregations)
.sort_values(["method", "group_label"])
.reset_index(drop=True)
)
subgroup_aggregate["coverage_std"] = subgroup_aggregate["coverage_std"].fillna(0.0)
subgroup_aggregate["coverage_gap_std"] = subgroup_aggregate["coverage_gap_std"].fillna(0.0)
if "average_set_size_std" in subgroup_aggregate.columns:
subgroup_aggregate["average_set_size_std"] = subgroup_aggregate["average_set_size_std"].fillna(0.0)
repeated_path = config.output_dir / "weighted_baselines_summary.csv"
subgroup_path = config.output_dir / "weighted_baselines_subgroups.csv"
config_path = config.output_dir / "config.json"
manifest_path = config.output_dir / "manifest.json"
repeated_filtered.to_csv(repeated_path, index=False)
subgroup_aggregate.to_csv(subgroup_path, index=False)
config_path.write_text(json.dumps(asdict(config), indent=2, default=str), encoding="utf-8")
manifest_path.write_text(
json.dumps(
{
"weighted_baselines_summary": str(repeated_path),
"weighted_baselines_subgroups": str(subgroup_path),
"config": str(config_path),
},
indent=2,
sort_keys=True,
),
encoding="utf-8",
)
return {
"weighted_baselines_summary": repeated_path,
"weighted_baselines_subgroups": subgroup_path,
"config": config_path,
"manifest": manifest_path,
}
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-weighted-baselines-analysis")
parser.add_argument("--sweep-dir", type=Path, required=True)
parser.add_argument("--model-type", default="xgboost")
parser.add_argument("--methods", type=_parse_methods, default=",".join(WEIGHTED_BASELINE_METHODS))
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_weighted_baselines_analysis(
WeightedBaselinesAnalysisConfig(
sweep_dir=args.sweep_dir,
output_dir=args.output_dir,
model_type=args.model_type,
methods=args.methods,
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())