zebra-open 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- zebra_open/__init__.py +33 -0
- zebra_open/cli.py +225 -0
- zebra_open/cohort.py +174 -0
- zebra_open/examples/README.md +36 -0
- zebra_open/examples/quickstart.ipynb +281 -0
- zebra_open/generator.py +798 -0
- zebra_open/io.py +31 -0
- zebra_open/model.py +101 -0
- zebra_open/retro.py +1315 -0
- zebra_open-0.1.0.dist-info/METADATA +447 -0
- zebra_open-0.1.0.dist-info/RECORD +15 -0
- zebra_open-0.1.0.dist-info/WHEEL +5 -0
- zebra_open-0.1.0.dist-info/entry_points.txt +6 -0
- zebra_open-0.1.0.dist-info/licenses/LICENSE +109 -0
- zebra_open-0.1.0.dist-info/top_level.txt +1 -0
zebra_open/__init__.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""zebra_open
|
|
2
|
+
|
|
3
|
+
Research-grade reference implementation of a ZeBRA-style retrospective risk modeling framework.
|
|
4
|
+
|
|
5
|
+
The CLI entry points are the recommended interface. Core functions are also
|
|
6
|
+
importable for scripting and integration.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from .generator import AgeSpec, GenParams, generate_dx_cohort, write_manifest, write_patients_jsonl
|
|
10
|
+
from .cohort import summarize_file, summarize_patients, summarize_aggregate
|
|
11
|
+
from .retro import (
|
|
12
|
+
train_model_from_files,
|
|
13
|
+
eval_retrospective_from_files,
|
|
14
|
+
score_patients_from_files,
|
|
15
|
+
load_model,
|
|
16
|
+
)
|
|
17
|
+
from .model import ZebraModel
|
|
18
|
+
|
|
19
|
+
__all__ = [
|
|
20
|
+
"AgeSpec",
|
|
21
|
+
"GenParams",
|
|
22
|
+
"generate_dx_cohort",
|
|
23
|
+
"write_patients_jsonl",
|
|
24
|
+
"write_manifest",
|
|
25
|
+
"summarize_file",
|
|
26
|
+
"summarize_patients",
|
|
27
|
+
"summarize_aggregate",
|
|
28
|
+
"train_model_from_files",
|
|
29
|
+
"eval_retrospective_from_files",
|
|
30
|
+
"score_patients_from_files",
|
|
31
|
+
"load_model",
|
|
32
|
+
"ZebraModel",
|
|
33
|
+
]
|
zebra_open/cli.py
ADDED
|
@@ -0,0 +1,225 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""zebra_open.cli
|
|
3
|
+
|
|
4
|
+
Console entry points:
|
|
5
|
+
zebra-generate : generate a synthetic DX-only cohort (research use)
|
|
6
|
+
zebra-cohort : summarize cohort (QC)
|
|
7
|
+
zebra-train : train a retrospective ZeBRA model
|
|
8
|
+
zebra-eval : retrospective evaluation (AUC + optional calibration)
|
|
9
|
+
zebra-score : prospective-style scoring (no labels)
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
import argparse
|
|
13
|
+
import json
|
|
14
|
+
import os
|
|
15
|
+
from pathlib import Path
|
|
16
|
+
|
|
17
|
+
from .model import ZebraModel
|
|
18
|
+
from . import retro
|
|
19
|
+
from .generator import GenParams, parse_yyyy_mm_dd, is_valid_icd10cm, generate_dx_cohort, write_patients_jsonl, write_manifest
|
|
20
|
+
from .cohort import summarize_file, summarize_aggregate
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def main_generate():
|
|
24
|
+
ap = argparse.ArgumentParser()
|
|
25
|
+
ap.add_argument("--out_dir", required=True, help="Output directory (will be created).")
|
|
26
|
+
ap.add_argument("--n_patients", type=int, default=10_000)
|
|
27
|
+
ap.add_argument("--target", required=True, help="Target ICD-10-CM code (e.g., J84.11).")
|
|
28
|
+
ap.add_argument("--case_frac", type=float, default=0.10)
|
|
29
|
+
ap.add_argument("--seed", type=int, default=13)
|
|
30
|
+
ap.add_argument("--start_date", default="2018-01-01", help="YYYY-MM-DD")
|
|
31
|
+
ap.add_argument("--end_date", default="2025-12-31", help="YYYY-MM-DD")
|
|
32
|
+
|
|
33
|
+
# Utilization controls (encounters/year)
|
|
34
|
+
ap.add_argument("--min_encounters_per_year", type=int, default=2)
|
|
35
|
+
ap.add_argument("--max_encounters_per_year", type=int, default=12)
|
|
36
|
+
|
|
37
|
+
# Chronic/problem list controls
|
|
38
|
+
ap.add_argument("--min_problem_list", type=int, default=0)
|
|
39
|
+
ap.add_argument("--max_problem_list", type=int, default=5)
|
|
40
|
+
|
|
41
|
+
# Target placement window (days from patient_start)
|
|
42
|
+
ap.add_argument("--min_days_to_target", type=int, default=30)
|
|
43
|
+
ap.add_argument("--max_days_to_target", type=int, default=2000)
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
# Target insertion controls
|
|
47
|
+
ap.add_argument("--target_min_occ", type=int, default=1)
|
|
48
|
+
ap.add_argument("--target_max_occ", type=int, default=3)
|
|
49
|
+
ap.add_argument("--target_last_third", action="store_true")
|
|
50
|
+
|
|
51
|
+
# Missingness controls
|
|
52
|
+
ap.add_argument("--missingness_encounter_drop", type=float, default=0.04)
|
|
53
|
+
ap.add_argument("--missingness_code_drop", type=float, default=0.06)
|
|
54
|
+
|
|
55
|
+
# Optional OpenAI enrichment
|
|
56
|
+
ap.add_argument("--enable_private_llm_pack", action="store_true")
|
|
57
|
+
ap.add_argument("--openai_model", default="gpt-5-nano")
|
|
58
|
+
ap.add_argument("--openai_timeout_s", type=int, default=30)
|
|
59
|
+
ap.add_argument("--openai_max_output_tokens", type=int, default=None)
|
|
60
|
+
|
|
61
|
+
args = ap.parse_args()
|
|
62
|
+
|
|
63
|
+
outdir = Path(args.out_dir)
|
|
64
|
+
outdir.mkdir(parents=True, exist_ok=True)
|
|
65
|
+
|
|
66
|
+
from .generator import GenParams, parse_yyyy_mm_dd, generate_dx_cohort, write_patients_jsonl, write_manifest, is_valid_icd10cm
|
|
67
|
+
|
|
68
|
+
target = str(args.target).strip().upper()
|
|
69
|
+
if not is_valid_icd10cm(target):
|
|
70
|
+
raise ValueError(f"--target {target!r} does not look like a valid ICD-10-CM code string.")
|
|
71
|
+
|
|
72
|
+
params = GenParams(
|
|
73
|
+
n_patients=int(args.n_patients),
|
|
74
|
+
case_frac=float(args.case_frac),
|
|
75
|
+
target_code=target,
|
|
76
|
+
start_date=parse_yyyy_mm_dd(args.start_date),
|
|
77
|
+
end_date=parse_yyyy_mm_dd(args.end_date),
|
|
78
|
+
seed=int(args.seed),
|
|
79
|
+
output=str(outdir / "patients.jsonl"),
|
|
80
|
+
jsonl=True,
|
|
81
|
+
target_min_occurrences=int(args.target_min_occ),
|
|
82
|
+
target_max_occurrences=int(args.target_max_occ),
|
|
83
|
+
target_last_third=bool(args.target_last_third),
|
|
84
|
+
missingness_encounter_drop=float(args.missingness_encounter_drop),
|
|
85
|
+
missingness_code_drop=float(args.missingness_code_drop),
|
|
86
|
+
min_encounters_per_year=int(args.min_encounters_per_year),
|
|
87
|
+
max_encounters_per_year=int(args.max_encounters_per_year),
|
|
88
|
+
min_problem_list=int(args.min_problem_list),
|
|
89
|
+
max_problem_list=int(args.max_problem_list),
|
|
90
|
+
min_days_to_target=int(args.min_days_to_target),
|
|
91
|
+
max_days_to_target=int(args.max_days_to_target),
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
patients, manifest = generate_dx_cohort(
|
|
95
|
+
params,
|
|
96
|
+
enable_openai_pack=bool(args.enable_private_llm_pack),
|
|
97
|
+
openai_model=str(args.openai_model),
|
|
98
|
+
openai_timeout_s=int(args.openai_timeout_s),
|
|
99
|
+
openai_max_output_tokens=(int(args.openai_max_output_tokens) if args.openai_max_output_tokens is not None else None),
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
patients_path = outdir / "patients.jsonl"
|
|
103
|
+
manifest_path = outdir / "cohort_manifest.json"
|
|
104
|
+
|
|
105
|
+
write_patients_jsonl(patients, str(patients_path))
|
|
106
|
+
write_manifest(manifest, str(manifest_path))
|
|
107
|
+
|
|
108
|
+
print(str(patients_path))
|
|
109
|
+
print(str(manifest_path))
|
|
110
|
+
return 0
|
|
111
|
+
|
|
112
|
+
def main_cohort():
|
|
113
|
+
ap = argparse.ArgumentParser()
|
|
114
|
+
ap.add_argument("--patients", required=True, help="Patients JSON or JSONL path.")
|
|
115
|
+
ap.add_argument("--target_prefix", required=True, help="Prefix to match (e.g., J84.11).")
|
|
116
|
+
ap.add_argument("--out", required=True, help="Output CSV (or .parquet) summary file.")
|
|
117
|
+
ap.add_argument("--aggregate_json", default=None, help="Optional aggregate summary JSON output.")
|
|
118
|
+
args = ap.parse_args()
|
|
119
|
+
|
|
120
|
+
df = summarize_file(args.patients, target_prefix=args.target_prefix, out_path=args.out)
|
|
121
|
+
agg = summarize_aggregate(df)
|
|
122
|
+
|
|
123
|
+
if args.aggregate_json:
|
|
124
|
+
with open(args.aggregate_json, "w", encoding="utf-8") as f:
|
|
125
|
+
json.dump(agg, f, indent=2)
|
|
126
|
+
|
|
127
|
+
print(args.out)
|
|
128
|
+
if args.aggregate_json:
|
|
129
|
+
print(args.aggregate_json)
|
|
130
|
+
return 0
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def main_train():
|
|
134
|
+
ap = argparse.ArgumentParser()
|
|
135
|
+
ap.add_argument("--patients", required=True)
|
|
136
|
+
ap.add_argument("--target", required=True, help="Comma-separated target codes or single code.")
|
|
137
|
+
ap.add_argument("--out", required=True, help="Output directory.")
|
|
138
|
+
ap.add_argument("--observation_days", type=int, default=365)
|
|
139
|
+
ap.add_argument("--horizon_days", type=int, default=28)
|
|
140
|
+
ap.add_argument("--prediction_days", type=int, default=365)
|
|
141
|
+
ap.add_argument("--confidence_days", type=int, default=365)
|
|
142
|
+
ap.add_argument("--or_infer_frac", type=float, default=0.60)
|
|
143
|
+
ap.add_argument("--val_frac", type=float, default=0.34)
|
|
144
|
+
ap.add_argument("--random_state", type=int, default=0)
|
|
145
|
+
args = ap.parse_args()
|
|
146
|
+
|
|
147
|
+
outdir = Path(args.out)
|
|
148
|
+
outdir.mkdir(parents=True, exist_ok=True)
|
|
149
|
+
|
|
150
|
+
targets = [t.strip() for t in args.target.split(",") if t.strip()]
|
|
151
|
+
|
|
152
|
+
m = ZebraModel.from_config(
|
|
153
|
+
observation_days=args.observation_days,
|
|
154
|
+
horizon_days=args.horizon_days,
|
|
155
|
+
prediction_days=args.prediction_days,
|
|
156
|
+
confidence_days=args.confidence_days,
|
|
157
|
+
or_infer_frac=args.or_infer_frac,
|
|
158
|
+
val_frac=args.val_frac,
|
|
159
|
+
random_state=args.random_state,
|
|
160
|
+
)
|
|
161
|
+
m.fit(patients_json=args.patients, target_codes=targets, out_dir=str(outdir))
|
|
162
|
+
return 0
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
def main_score():
|
|
166
|
+
ap = argparse.ArgumentParser()
|
|
167
|
+
ap.add_argument("--model", required=True, help="Path to zebra_retro_model.joblib.")
|
|
168
|
+
ap.add_argument("--patients", required=True, help="Path to patients JSON or JSONL.")
|
|
169
|
+
ap.add_argument("--out", required=True, help="Output CSV path.")
|
|
170
|
+
ap.add_argument("--as_of", default=None, help="Optional YYYY-MM-DD as-of date for all patients.")
|
|
171
|
+
args = ap.parse_args()
|
|
172
|
+
|
|
173
|
+
m = ZebraModel.load(args.model)
|
|
174
|
+
m.predict_proba(args.patients, as_of=args.as_of, out_csv=args.out)
|
|
175
|
+
return 0
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
def main_eval():
|
|
179
|
+
ap = argparse.ArgumentParser()
|
|
180
|
+
ap.add_argument("--model", required=True, help="Path to zebra_retro_model.joblib.")
|
|
181
|
+
ap.add_argument("--patients", required=True, help="Path to patients JSON or JSONL to evaluate.")
|
|
182
|
+
ap.add_argument("--out", required=True, help="Output CSV path with columns patient_id,t0,y,p.")
|
|
183
|
+
ap.add_argument("--summary", default=None, help="Optional JSON summary path (includes AUC when computable).")
|
|
184
|
+
ap.add_argument("--calibration_bins", type=int, default=10, help="Number of bins for calibration curve.")
|
|
185
|
+
ap.add_argument("--calibration_csv", default=None, help="Optional CSV path for binned calibration stats.")
|
|
186
|
+
ap.add_argument("--calibration_plot", default=None, help="Optional PNG path for a reliability diagram.")
|
|
187
|
+
ap.add_argument("--target", default=None, help="Optional target code(s) override. Comma-separated or single.")
|
|
188
|
+
ap.add_argument("--max_controls", type=int, default=None, help="Optional control downsampling for evaluation.")
|
|
189
|
+
args = ap.parse_args()
|
|
190
|
+
|
|
191
|
+
target_codes = None
|
|
192
|
+
if args.target:
|
|
193
|
+
target_codes = [t.strip() for t in args.target.split(",") if t.strip()]
|
|
194
|
+
|
|
195
|
+
# If --out is a directory, write default artifacts inside it.
|
|
196
|
+
out_csv = args.out
|
|
197
|
+
out_is_dir = out_csv.endswith(os.sep) or (os.path.exists(out_csv) and os.path.isdir(out_csv))
|
|
198
|
+
if out_is_dir:
|
|
199
|
+
os.makedirs(out_csv, exist_ok=True)
|
|
200
|
+
out_csv = os.path.join(out_csv, 'eval.csv')
|
|
201
|
+
|
|
202
|
+
summary_path = args.summary
|
|
203
|
+
if summary_path is None and out_is_dir:
|
|
204
|
+
summary_path = os.path.join(args.out, 'eval_summary.json')
|
|
205
|
+
|
|
206
|
+
cal_csv = args.calibration_csv
|
|
207
|
+
cal_plot = args.calibration_plot
|
|
208
|
+
if out_is_dir:
|
|
209
|
+
if cal_csv is None:
|
|
210
|
+
cal_csv = os.path.join(args.out, 'calibration.csv')
|
|
211
|
+
if cal_plot is None:
|
|
212
|
+
cal_plot = os.path.join(args.out, 'calibration.png')
|
|
213
|
+
|
|
214
|
+
retro.eval_retrospective_from_files(
|
|
215
|
+
model_path=args.model,
|
|
216
|
+
patients_json=args.patients,
|
|
217
|
+
target_codes=target_codes,
|
|
218
|
+
out_csv=out_csv,
|
|
219
|
+
out_summary_json=summary_path,
|
|
220
|
+
calibration_bins=int(args.calibration_bins),
|
|
221
|
+
calibration_csv=cal_csv,
|
|
222
|
+
calibration_plot=cal_plot,
|
|
223
|
+
max_controls=args.max_controls,
|
|
224
|
+
)
|
|
225
|
+
return 0
|
zebra_open/cohort.py
ADDED
|
@@ -0,0 +1,174 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""
|
|
3
|
+
zebra_open.cohort
|
|
4
|
+
|
|
5
|
+
Cohort summarization utilities.
|
|
6
|
+
|
|
7
|
+
This module is derived from code/cohort_visualizer.py and provides reusable
|
|
8
|
+
functions to summarize patient DX history and target-prefix presence.
|
|
9
|
+
|
|
10
|
+
It does not perform ZeBRA retrospective window labeling. It is intended for
|
|
11
|
+
data sanity checks, dataset characterization, and quick QC before model runs.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
from datetime import date, datetime
|
|
17
|
+
from typing import Any, Dict, Iterable, List, Optional, Tuple
|
|
18
|
+
import json
|
|
19
|
+
|
|
20
|
+
import pandas as pd
|
|
21
|
+
|
|
22
|
+
DATE_FMT = "%m/%d/%Y"
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def parse_mmddyyyy(s: str) -> date:
|
|
26
|
+
return datetime.strptime(s, DATE_FMT).date()
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def norm_code(code: str) -> str:
|
|
30
|
+
return (code or "").strip().upper()
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def code_matches_prefix(code: str, target_prefix: str) -> bool:
|
|
34
|
+
c = norm_code(code)
|
|
35
|
+
t = norm_code(target_prefix)
|
|
36
|
+
return c.startswith(t) if t else False
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def load_patients(path: str) -> List[Dict[str, Any]]:
|
|
40
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
41
|
+
first = f.read(1)
|
|
42
|
+
f.seek(0)
|
|
43
|
+
if first == "[":
|
|
44
|
+
return json.load(f)
|
|
45
|
+
out = []
|
|
46
|
+
for line in f:
|
|
47
|
+
line = line.strip()
|
|
48
|
+
if not line:
|
|
49
|
+
continue
|
|
50
|
+
out.append(json.loads(line))
|
|
51
|
+
return out
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def extract_dx_records(p: Dict[str, Any]) -> List[Tuple[date, str]]:
|
|
55
|
+
dx = p.get("DX_record", []) or []
|
|
56
|
+
out: List[Tuple[date, str]] = []
|
|
57
|
+
for r in dx:
|
|
58
|
+
d = r.get("date")
|
|
59
|
+
c = r.get("code")
|
|
60
|
+
if not d or not c:
|
|
61
|
+
continue
|
|
62
|
+
try:
|
|
63
|
+
out.append((parse_mmddyyyy(d), norm_code(c)))
|
|
64
|
+
except Exception:
|
|
65
|
+
continue
|
|
66
|
+
out.sort(key=lambda t: t[0])
|
|
67
|
+
return out
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def summarize_patient(p: Dict[str, Any], target_prefix: str) -> Dict[str, Any]:
|
|
71
|
+
pid = p.get("patient_id", p.get("pid", p.get("id", None)))
|
|
72
|
+
pid = "" if pid is None else str(pid)
|
|
73
|
+
|
|
74
|
+
dx = extract_dx_records(p)
|
|
75
|
+
n_dx = len(dx)
|
|
76
|
+
uniq = len(set(c for _, c in dx))
|
|
77
|
+
|
|
78
|
+
if n_dx == 0:
|
|
79
|
+
return {
|
|
80
|
+
"patient_id": pid,
|
|
81
|
+
"target_prefix": target_prefix,
|
|
82
|
+
"target": 0,
|
|
83
|
+
"n_dx_records": 0,
|
|
84
|
+
"n_unique_dx_codes": 0,
|
|
85
|
+
"first_observed_date": None,
|
|
86
|
+
"last_observed_date": None,
|
|
87
|
+
"duration_days": None,
|
|
88
|
+
"first_target_date": None,
|
|
89
|
+
"days_before_first_target": None,
|
|
90
|
+
"days_after_first_target": None,
|
|
91
|
+
}
|
|
92
|
+
|
|
93
|
+
first_d = dx[0][0]
|
|
94
|
+
last_d = dx[-1][0]
|
|
95
|
+
duration_days = (last_d - first_d).days
|
|
96
|
+
|
|
97
|
+
first_t: Optional[date] = None
|
|
98
|
+
for d, c in dx:
|
|
99
|
+
if code_matches_prefix(c, target_prefix):
|
|
100
|
+
first_t = d
|
|
101
|
+
break
|
|
102
|
+
|
|
103
|
+
is_case = 1 if first_t is not None else 0
|
|
104
|
+
|
|
105
|
+
if first_t is None:
|
|
106
|
+
return {
|
|
107
|
+
"patient_id": pid,
|
|
108
|
+
"target_prefix": target_prefix,
|
|
109
|
+
"target": 0,
|
|
110
|
+
"n_dx_records": n_dx,
|
|
111
|
+
"n_unique_dx_codes": uniq,
|
|
112
|
+
"first_observed_date": first_d.strftime(DATE_FMT),
|
|
113
|
+
"last_observed_date": last_d.strftime(DATE_FMT),
|
|
114
|
+
"duration_days": duration_days,
|
|
115
|
+
"first_target_date": None,
|
|
116
|
+
"days_before_first_target": None,
|
|
117
|
+
"days_after_first_target": None,
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
return {
|
|
121
|
+
"patient_id": pid,
|
|
122
|
+
"target_prefix": target_prefix,
|
|
123
|
+
"target": is_case,
|
|
124
|
+
"n_dx_records": n_dx,
|
|
125
|
+
"n_unique_dx_codes": uniq,
|
|
126
|
+
"first_observed_date": first_d.strftime(DATE_FMT),
|
|
127
|
+
"last_observed_date": last_d.strftime(DATE_FMT),
|
|
128
|
+
"duration_days": duration_days,
|
|
129
|
+
"first_target_date": first_t.strftime(DATE_FMT),
|
|
130
|
+
"days_before_first_target": (first_t - first_d).days,
|
|
131
|
+
"days_after_first_target": (last_d - first_t).days,
|
|
132
|
+
}
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def summarize_patients(patients: Iterable[Dict[str, Any]], target_prefix: str) -> pd.DataFrame:
|
|
136
|
+
rows = [summarize_patient(p, target_prefix=target_prefix) for p in patients]
|
|
137
|
+
return pd.DataFrame(rows)
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def summarize_file(in_path: str, target_prefix: str, out_path: Optional[str] = None) -> pd.DataFrame:
|
|
141
|
+
pats = load_patients(in_path)
|
|
142
|
+
df = summarize_patients(pats, target_prefix=target_prefix)
|
|
143
|
+
if out_path:
|
|
144
|
+
if out_path.endswith(".parquet"):
|
|
145
|
+
df.to_parquet(out_path, index=False)
|
|
146
|
+
else:
|
|
147
|
+
df.to_csv(out_path, index=False)
|
|
148
|
+
return df
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def summarize_aggregate(df: pd.DataFrame) -> Dict[str, Any]:
|
|
152
|
+
# Basic aggregate stats for quick QC.
|
|
153
|
+
n = int(df.shape[0])
|
|
154
|
+
n_cases = int(df["target"].sum()) if "target" in df.columns else 0
|
|
155
|
+
n_controls = int(n - n_cases)
|
|
156
|
+
out: Dict[str, Any] = {"n": n, "cases": n_cases, "controls": n_controls}
|
|
157
|
+
|
|
158
|
+
if "duration_days" in df.columns:
|
|
159
|
+
d = pd.to_numeric(df["duration_days"], errors="coerce")
|
|
160
|
+
out["duration_days"] = {
|
|
161
|
+
"n_finite": int(d.notna().sum()),
|
|
162
|
+
"min": float(d.min()) if d.notna().any() else None,
|
|
163
|
+
"median": float(d.median()) if d.notna().any() else None,
|
|
164
|
+
"max": float(d.max()) if d.notna().any() else None,
|
|
165
|
+
}
|
|
166
|
+
|
|
167
|
+
if "n_dx_records" in df.columns:
|
|
168
|
+
x = pd.to_numeric(df["n_dx_records"], errors="coerce")
|
|
169
|
+
out["n_dx_records"] = {
|
|
170
|
+
"min": float(x.min()) if x.notna().any() else None,
|
|
171
|
+
"median": float(x.median()) if x.notna().any() else None,
|
|
172
|
+
"max": float(x.max()) if x.notna().any() else None,
|
|
173
|
+
}
|
|
174
|
+
return out
|
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
# ZeBRA Open examples
|
|
2
|
+
|
|
3
|
+
This directory is distributed with the `zebra-open` Python package.
|
|
4
|
+
|
|
5
|
+
## Quickstart notebook
|
|
6
|
+
|
|
7
|
+
`quickstart.ipynb` demonstrates an end-to-end ZeBRA Open workflow using only synthetic data:
|
|
8
|
+
|
|
9
|
+
1. generate a longitudinal synthetic cohort;
|
|
10
|
+
2. inspect the generated cohort;
|
|
11
|
+
3. train a retrospective ZeBRA model;
|
|
12
|
+
4. evaluate discrimination and calibration; and
|
|
13
|
+
5. score patients with the fitted model.
|
|
14
|
+
|
|
15
|
+
No clinical or private data are required.
|
|
16
|
+
|
|
17
|
+
When working from a clone of the repository, install the package first:
|
|
18
|
+
|
|
19
|
+
```bash
|
|
20
|
+
python3 -m pip install -e ./python
|
|
21
|
+
```
|
|
22
|
+
|
|
23
|
+
After a PyPI release, the corresponding command is:
|
|
24
|
+
|
|
25
|
+
```bash
|
|
26
|
+
python3 -m pip install zebra-open
|
|
27
|
+
```
|
|
28
|
+
|
|
29
|
+
The installed notebook can also be located programmatically:
|
|
30
|
+
|
|
31
|
+
```python
|
|
32
|
+
from importlib.resources import files
|
|
33
|
+
|
|
34
|
+
notebook = files("zebra_open").joinpath("examples", "quickstart.ipynb")
|
|
35
|
+
print(notebook)
|
|
36
|
+
```
|