from __future__ import annotations import argparse import json from dataclasses import asdict, dataclass from pathlib import Path from time import perf_counter import pandas as pd from sepsis_mcp.conformal import CPMDAExactClassifier, MissingnessAwareConformalClassifier, SplitConformalClassifier from sepsis_mcp.gossis import build_hospital_disjoint_split_with_selection, load_gossis_dataset from sepsis_mcp.gossis_experiment import ( GossisRunConfig, _build_structured_grouping, _encode_feature_splits, _missingness_mask_matrix, ) from sepsis_mcp.modeling import ProbabilityEstimator @dataclass class RuntimeAnalysisConfig: data_root: Path output_dir: Path model_type: str = "xgboost" random_state: int = 0 alpha: float = 0.1 selection_fraction: float = 0.1 min_hospital_admissions: int = 500 min_selection_group_rows: int = 100 def summarize_runtime_records(records: pd.DataFrame) -> pd.DataFrame: if records.empty: return pd.DataFrame() summary = records.copy() standard = ( summary[summary["method"] == "standard"][["stage", "seconds"]] .rename(columns={"seconds": "standard_seconds"}) ) summary = summary.merge(standard, on="stage", how="left", validate="many_to_one") summary["relative_to_standard"] = (summary["seconds"] / summary["standard_seconds"]).round(6) summary.loc[summary["stage"] == "variable_selection", "relative_to_standard"] = pd.NA return summary.sort_values(["stage", "method"]).reset_index(drop=True) def run_runtime_analysis(config: RuntimeAnalysisConfig) -> dict[str, Path]: config.output_dir.mkdir(parents=True, exist_ok=True) dataset = load_gossis_dataset(config.data_root, min_hospital_admissions=config.min_hospital_admissions) split = build_hospital_disjoint_split_with_selection( dataset.frame, train_fraction=0.6, selection_fraction=config.selection_fraction, calibration_fraction=0.1, random_state=config.random_state, ) train_features, calibration_features, test_features = _encode_feature_splits( split.train_frame, split.calibration_frame, split.test_frame, dataset.feature_columns, ) selection_features, _, _ = _encode_feature_splits( split.train_frame, split.selection_frame, split.test_frame, dataset.feature_columns, ) estimator = ProbabilityEstimator(random_state=config.random_state, model_type=config.model_type) estimator.fit(train_features, split.train_frame["label"]) selection_probabilities = estimator.predict_positive_proba(selection_features) calibration_probabilities = estimator.predict_positive_proba(calibration_features) test_probabilities = estimator.predict_positive_proba(test_features) records: list[dict[str, float | int | str]] = [] selection_start = perf_counter() structured_grouping = _build_structured_grouping( config=GossisRunConfig( data_root=config.data_root, alpha=config.alpha, selection_fraction=config.selection_fraction, model_type=config.model_type, random_state=config.random_state, missingness_grouping_strategy="coverage_gap_variable", min_selection_group_rows=config.min_selection_group_rows, ), feature_columns=dataset.feature_columns, selection_frame=split.selection_frame, calibration_frame=split.calibration_frame, test_frame=split.test_frame, selection_probabilities=selection_probabilities, calibration_probabilities=calibration_probabilities, test_probabilities=test_probabilities, ) records.append({"method": "missingness_grouping", "stage": "variable_selection", "seconds": perf_counter() - selection_start}) method_objects = { "standard": SplitConformalClassifier(alpha=config.alpha), "missingness_aware": MissingnessAwareConformalClassifier( alpha=config.alpha, min_group_size=max(2, min(10, len(split.calibration_frame) // 5)), ), "cp_mda_exact": CPMDAExactClassifier( alpha=config.alpha, top_k_features=10, min_match=10, ), } calibration_payloads = { "standard": { "calibration_labels": split.calibration_frame["label"].tolist(), "calibration_positive_probabilities": calibration_probabilities.tolist(), }, "missingness_aware": { "calibration_labels": split.calibration_frame["label"].tolist(), "calibration_positive_probabilities": calibration_probabilities.tolist(), "calibration_group_ids": structured_grouping.calibration_groups["group"].tolist(), }, "cp_mda_exact": { "calibration_labels": split.calibration_frame["label"].tolist(), "calibration_positive_probabilities": calibration_probabilities.tolist(), "calibration_masks": _missingness_mask_matrix(split.calibration_frame, dataset.feature_columns), "feature_names": dataset.feature_columns, }, } prediction_payloads = { "standard": {"positive_probabilities": test_probabilities.tolist()}, "missingness_aware": { "positive_probabilities": test_probabilities.tolist(), "test_group_ids": structured_grouping.test_groups["group"].tolist(), }, "cp_mda_exact": { "positive_probabilities": test_probabilities.tolist(), "test_masks": _missingness_mask_matrix(split.test_frame, dataset.feature_columns), }, } for method_name, method in method_objects.items(): start = perf_counter() method.fit(**calibration_payloads[method_name]) records.append({"method": method_name, "stage": "calibration", "seconds": perf_counter() - start}) start = perf_counter() method.predict_sets(**prediction_payloads[method_name]) records.append( { "method": method_name, "stage": "test", "seconds": perf_counter() - start, } ) records_frame = pd.DataFrame(records) summary = summarize_runtime_records(records_frame) records_path = config.output_dir / "runtime_records.csv" summary_path = config.output_dir / "runtime_summary.csv" config_path = config.output_dir / "config.json" manifest_path = config.output_dir / "manifest.json" records_frame.to_csv(records_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( { "runtime_records": str(records_path), "runtime_summary": str(summary_path), "config": str(config_path), }, indent=2, sort_keys=True, ), encoding="utf-8", ) return { "runtime_records": records_path, "runtime_summary": summary_path, "config": config_path, "manifest": manifest_path, } def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="sepsis-mcp-appendix-runtime-analysis") parser.add_argument("--data-root", type=Path, required=True) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--model-type", default="xgboost") parser.add_argument("--random-state", type=int, default=0) 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_runtime_analysis( RuntimeAnalysisConfig( data_root=args.data_root, output_dir=args.output_dir, model_type=args.model_type, random_state=args.random_state, 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())