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