File size: 4,066 Bytes
32f5a65
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
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())