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