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