misscp / src /sepsis_mcp /cli.py
Anonymous
Initial anonymous MissCP release
32f5a65
Raw
History Blame Contribute Delete
4.07 kB
from __future__ import annotations
import argparse
from pathlib import Path
from sepsis_mcp.runner import RunConfig, run_experiments
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(prog="sepsis-mcp")
subparsers = parser.add_subparsers(dest="command", required=True)
run_parser = subparsers.add_parser("run")
run_parser.add_argument("--data-root", type=Path, required=True)
run_parser.add_argument("--train-hospital", choices=["A", "B"], default="A")
run_parser.add_argument("--test-mode", choices=["aa", "ab", "both"], default="both")
run_parser.add_argument(
"--model-type",
choices=["sklearn_gbdt", "xgboost", "logistic_regression"],
default="sklearn_gbdt",
)
run_parser.add_argument("--train-patients", type=int, default=128)
run_parser.add_argument("--selection-patients", type=int, default=0)
run_parser.add_argument("--calibration-patients", type=int, default=64)
run_parser.add_argument("--test-patients", type=int, default=128)
run_parser.add_argument("--lookback-hours", type=int, default=6)
run_parser.add_argument("--horizon-hours", type=int, default=6)
run_parser.add_argument("--alpha", type=float, default=0.1)
run_parser.add_argument("--random-state", type=int, default=0)
run_parser.add_argument("--weighted-shrinkage-lambda", type=float, default=0.5)
run_parser.add_argument("--mondrian-tilted-shrinkage-lambda", type=float, default=0.5)
run_parser.add_argument("--mondrian-tilted-bandwidth", type=float, default=None)
run_parser.add_argument(
"--missingness-grouping-strategy",
choices=["coverage_gap_variable", "mask_cluster"],
default=None,
)
run_parser.add_argument("--enable-external-baselines", action="store_true")
run_parser.add_argument("--enable-learned-partition", action="store_true")
run_parser.add_argument(
"--mask-strategy",
choices=["none", "random_drop", "block_missing", "selective_missing"],
default="none",
)
run_parser.add_argument("--mask-rate", type=float, default=0.3)
run_parser.add_argument("--block-length", type=int, default=4)
run_parser.add_argument("--mask-random-state", type=int, default=0)
run_parser.add_argument("--selective-feature", action="append", dest="selective_features")
run_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)
if args.command == "run":
config = RunConfig(
data_root=args.data_root,
train_hospital=args.train_hospital,
test_mode=args.test_mode,
model_type=args.model_type,
train_patients=args.train_patients,
selection_patients=args.selection_patients,
calibration_patients=args.calibration_patients,
test_patients=args.test_patients,
lookback_hours=args.lookback_hours,
horizon_hours=args.horizon_hours,
alpha=args.alpha,
random_state=args.random_state,
weighted_shrinkage_lambda=args.weighted_shrinkage_lambda,
mondrian_tilted_shrinkage_lambda=args.mondrian_tilted_shrinkage_lambda,
mondrian_tilted_bandwidth=args.mondrian_tilted_bandwidth,
missingness_grouping_strategy=args.missingness_grouping_strategy,
enable_external_baselines=args.enable_external_baselines,
enable_learned_partition=args.enable_learned_partition,
mask_strategy=args.mask_strategy,
mask_rate=args.mask_rate,
block_length=args.block_length,
mask_random_state=args.mask_random_state,
selective_features=args.selective_features,
output_dir=args.output_dir,
)
run_experiments(config)
return 0
parser.error(f"unsupported command: {args.command}")
return 2
if __name__ == "__main__":
raise SystemExit(main())