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 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
+ ```