misscp / src /sepsis_mcp /appendix_runtime_analysis.py
Anonymous
Initial anonymous MissCP release
32f5a65
Raw
History Blame Contribute Delete
8.54 kB
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())