electrotrace 1.9.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.
Files changed (42) hide show
  1. electrotrace/__init__.py +28 -0
  2. electrotrace/__main__.py +3 -0
  3. electrotrace/annotations.py +194 -0
  4. electrotrace/baseline_detectors.py +122 -0
  5. electrotrace/beats.py +47 -0
  6. electrotrace/benchmark.py +154 -0
  7. electrotrace/candidate_suppressor.py +441 -0
  8. electrotrace/cli.py +630 -0
  9. electrotrace/detectors.py +152 -0
  10. electrotrace/formats.py +147 -0
  11. electrotrace/fp_analysis.py +358 -0
  12. electrotrace/hearttwin_adapter.py +90 -0
  13. electrotrace/io.py +163 -0
  14. electrotrace/lead_quality.py +138 -0
  15. electrotrace/lead_selection.py +169 -0
  16. electrotrace/metadata.py +111 -0
  17. electrotrace/ml.py +178 -0
  18. electrotrace/phenotype.py +84 -0
  19. electrotrace/phenotype_validation.py +133 -0
  20. electrotrace/polarity_v2.py +201 -0
  21. electrotrace/project.py +43 -0
  22. electrotrace/project_store.py +147 -0
  23. electrotrace/provenance.py +155 -0
  24. electrotrace/qrs_delineation.py +139 -0
  25. electrotrace/qtdb_detector_adapter.py +7 -0
  26. electrotrace/research_validation.py +146 -0
  27. electrotrace/scale_estimation.py +172 -0
  28. electrotrace/security.py +57 -0
  29. electrotrace/server_app.py +375 -0
  30. electrotrace/signal.py +84 -0
  31. electrotrace/statistics.py +98 -0
  32. electrotrace/threshold_selection.py +40 -0
  33. electrotrace/validation.py +241 -0
  34. electrotrace/validation_detectors.py +570 -0
  35. electrotrace/wfdb_records.py +375 -0
  36. electrotrace/window.py +117 -0
  37. electrotrace-1.9.0.dist-info/METADATA +187 -0
  38. electrotrace-1.9.0.dist-info/RECORD +42 -0
  39. electrotrace-1.9.0.dist-info/WHEEL +5 -0
  40. electrotrace-1.9.0.dist-info/entry_points.txt +3 -0
  41. electrotrace-1.9.0.dist-info/licenses/LICENSE +20 -0
  42. electrotrace-1.9.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,28 @@
1
+ """ElectroTrace: reproducible ECG/electrophysiology research and benchmarking tools."""
2
+
3
+ __version__ = "1.9.0"
4
+
5
+ from .candidate_suppressor import CandidateSuppressor
6
+ from .io import load_recording
7
+ from .provenance import DatasetManifest, manifest_from_dict
8
+ from .qrs_delineation import delineate_qrs
9
+ from .research_validation import build_validation_report, summarize_records_rigorous, write_validation_report
10
+ from .signal import apply_pipeline
11
+ from .validation import validate_record
12
+ from .validation_detectors import detect_r_peaks, detect_r_peaks_two_stage
13
+
14
+ __all__ = [
15
+ "__version__",
16
+ "CandidateSuppressor",
17
+ "DatasetManifest",
18
+ "manifest_from_dict",
19
+ "load_recording",
20
+ "apply_pipeline",
21
+ "detect_r_peaks",
22
+ "detect_r_peaks_two_stage",
23
+ "delineate_qrs",
24
+ "validate_record",
25
+ "build_validation_report",
26
+ "summarize_records_rigorous",
27
+ "write_validation_report",
28
+ ]
@@ -0,0 +1,3 @@
1
+ from .cli import main
2
+
3
+ raise SystemExit(main())
@@ -0,0 +1,194 @@
1
+ """Annotation model, validation, serialization, review state, and agreement metrics."""
2
+ from __future__ import annotations
3
+
4
+ import json
5
+ import math
6
+ import uuid
7
+ from dataclasses import asdict, dataclass, field
8
+ from typing import Literal, Optional
9
+
10
+ AnnotationType = Literal["interval", "point"]
11
+ ReviewStatus = Literal["unreviewed", "accepted", "flagged"]
12
+ DEFAULT_LABELS = ["P Wave", "QRS", "T Wave", "Pacemaker Spike", "R Peak", "Artifact", "Abnormal Beat", "Graft Activation", "Arrhythmia", "Custom"]
13
+
14
+
15
+ @dataclass
16
+ class Annotation:
17
+ label: str
18
+ type: AnnotationType
19
+ channel: str
20
+ start: Optional[float] = None
21
+ end: Optional[float] = None
22
+ time: Optional[float] = None
23
+ confidence: float = 1.0
24
+ notes: str = ""
25
+ annotator: str = ""
26
+ status: ReviewStatus = "unreviewed"
27
+ reviewer: str = ""
28
+ review_notes: str = ""
29
+ id: str = field(default_factory=lambda: uuid.uuid4().hex[:10])
30
+
31
+ def validate(self, duration_s: float | None = None, start_time_s: float = 0.0, end_time_s: float | None = None) -> None:
32
+ if not self.label.strip():
33
+ raise ValueError("label must not be empty")
34
+ if not self.channel.strip():
35
+ raise ValueError("channel must not be empty")
36
+ if not math.isfinite(float(self.confidence)) or not 0.0 <= self.confidence <= 1.0:
37
+ raise ValueError("confidence must be between 0 and 1")
38
+ if self.status not in {"unreviewed", "accepted", "flagged"}:
39
+ raise ValueError("invalid review status")
40
+ start_bound = float(start_time_s)
41
+ end_bound = float(end_time_s) if end_time_s is not None else (start_bound + float(duration_s) if duration_s is not None else None)
42
+ if not math.isfinite(start_bound) or (end_bound is not None and not math.isfinite(end_bound)) or (end_bound is not None and end_bound < start_bound):
43
+ raise ValueError("invalid recording time bounds")
44
+ if self.type == "interval":
45
+ if self.start is None or self.end is None:
46
+ raise ValueError("interval annotations require start and end")
47
+ if not math.isfinite(self.start) or not math.isfinite(self.end):
48
+ raise ValueError("interval coordinates must be finite")
49
+ if self.end <= self.start:
50
+ raise ValueError("end must be greater than start")
51
+ if self.time is not None:
52
+ raise ValueError("interval annotations must not define time")
53
+ if end_bound is not None and (self.start < start_bound or self.end > end_bound):
54
+ raise ValueError("interval is outside the recording bounds")
55
+ elif self.type == "point":
56
+ if self.time is None:
57
+ raise ValueError("point annotations require time")
58
+ if not math.isfinite(self.time):
59
+ raise ValueError("point time must be finite")
60
+ if self.start is not None or self.end is not None:
61
+ raise ValueError("point annotations must not define start/end")
62
+ if end_bound is not None and not start_bound <= self.time <= end_bound:
63
+ raise ValueError("point is outside the recording bounds")
64
+ else:
65
+ raise ValueError(f"unknown annotation type: {self.type}")
66
+
67
+ @property
68
+ def position(self) -> float:
69
+ return self.start if self.type == "interval" else self.time # type: ignore[return-value]
70
+
71
+ def to_dict(self) -> dict:
72
+ return asdict(self)
73
+
74
+
75
+ class AnnotationStore:
76
+ def __init__(self, duration_s: float | None = None, start_time_s: float = 0.0, end_time_s: float | None = None):
77
+ self.duration_s = duration_s
78
+ self.start_time_s = float(start_time_s)
79
+ self.end_time_s = float(end_time_s) if end_time_s is not None else (self.start_time_s + float(duration_s) if duration_s is not None else None)
80
+ self._items: list[Annotation] = []
81
+
82
+ @property
83
+ def items(self) -> list[Annotation]:
84
+ return list(self._items)
85
+
86
+ def add(self, ann: Annotation) -> Annotation:
87
+ ann.validate(self.duration_s, self.start_time_s, self.end_time_s)
88
+ if any(a.id == ann.id for a in self._items):
89
+ raise ValueError(f"duplicate annotation id: {ann.id}")
90
+ self._items.append(ann)
91
+ self._sort()
92
+ return ann
93
+
94
+ def update(self, ann_id: str, **changes) -> Annotation:
95
+ for idx, existing in enumerate(self._items):
96
+ if existing.id == ann_id:
97
+ data = {**asdict(existing), **changes}
98
+ updated = Annotation(**data)
99
+ updated.validate(self.duration_s, self.start_time_s, self.end_time_s)
100
+ self._items[idx] = updated
101
+ self._sort()
102
+ return updated
103
+ raise KeyError(ann_id)
104
+
105
+ def delete(self, ann_id: str) -> bool:
106
+ before = len(self._items)
107
+ self._items = [a for a in self._items if a.id != ann_id]
108
+ return len(self._items) != before
109
+
110
+ def duplicate(self, ann_id: str) -> Annotation:
111
+ for existing in self._items:
112
+ if existing.id == ann_id:
113
+ data = asdict(existing)
114
+ data["id"] = uuid.uuid4().hex[:10]
115
+ data["status"] = "unreviewed"
116
+ data["reviewer"] = ""
117
+ data["review_notes"] = ""
118
+ return self.add(Annotation(**data))
119
+ raise KeyError(ann_id)
120
+
121
+ def clear(self) -> None:
122
+ self._items.clear()
123
+
124
+ def _sort(self) -> None:
125
+ self._items.sort(key=lambda a: a.position)
126
+
127
+ def to_dict(self, source_file: str = "", metadata: dict | None = None) -> dict:
128
+ meta = dict(metadata or {})
129
+ meta.setdefault("time_start_s", self.start_time_s)
130
+ if self.end_time_s is not None:
131
+ meta.setdefault("time_end_s", self.end_time_s)
132
+ return {"schema": "electrotrace.annotation/v2", "file": source_file, "metadata": meta, "annotations": [a.to_dict() for a in self._items]}
133
+
134
+ def to_json(self, source_file: str = "", metadata: dict | None = None) -> str:
135
+ return json.dumps(self.to_dict(source_file, metadata), indent=2)
136
+
137
+ @classmethod
138
+ def from_dict(cls, data: dict, duration_s: float | None = None, start_time_s: float | None = None, end_time_s: float | None = None) -> "AnnotationStore":
139
+ metadata = data.get("metadata") or {}
140
+ start = float(metadata.get("time_start_s", 0.0) if start_time_s is None else start_time_s)
141
+ end = metadata.get("time_end_s") if end_time_s is None else end_time_s
142
+ store = cls(duration_s=duration_s, start_time_s=start, end_time_s=float(end) if end is not None else None)
143
+ schema = data.get("schema", data.get("annotation_schema", ""))
144
+ if schema and schema not in {"electrotrace.annotation/v2", "v1"}:
145
+ raise ValueError(f"unsupported annotation schema: {schema}")
146
+ for raw in data.get("annotations", []):
147
+ store.add(Annotation(**raw))
148
+ return store
149
+
150
+ @classmethod
151
+ def from_json(cls, text: str, duration_s: float | None = None, start_time_s: float | None = None, end_time_s: float | None = None) -> "AnnotationStore":
152
+ try:
153
+ data = json.loads(text)
154
+ except json.JSONDecodeError as exc:
155
+ raise ValueError(f"invalid JSON: {exc}") from exc
156
+ return cls.from_dict(data, duration_s=duration_s, start_time_s=start_time_s, end_time_s=end_time_s)
157
+
158
+ def to_csv_rows(self, source_file: str = "") -> list[dict]:
159
+ return [{"file": source_file, "id": a.id, "type": a.type, "label": a.label, "channel": a.channel, "start": a.start, "end": a.end, "time": a.time, "confidence": a.confidence, "notes": a.notes, "annotator": a.annotator, "status": a.status, "reviewer": a.reviewer, "review_notes": a.review_notes} for a in self._items]
160
+
161
+
162
+ def point_agreement(a: list[Annotation], b: list[Annotation], tolerance_s: float = 0.04) -> dict:
163
+ if tolerance_s <= 0 or not math.isfinite(tolerance_s):
164
+ raise ValueError("tolerance_s must be positive and finite")
165
+ left = [x for x in a if x.type == "point"]
166
+ right = [x for x in b if x.type == "point"]
167
+ used: set[str] = set()
168
+ errors = []
169
+ matches = 0
170
+ for x in left:
171
+ candidates = [
172
+ y for y in right
173
+ if y.id not in used
174
+ and y.label == x.label
175
+ and y.channel == x.channel
176
+ and y.time is not None
177
+ and x.time is not None
178
+ and abs(y.time - x.time) <= tolerance_s
179
+ ]
180
+ if candidates:
181
+ y = min(candidates, key=lambda z: abs(z.time - x.time))
182
+ used.add(y.id)
183
+ matches += 1
184
+ errors.append(abs(y.time - x.time))
185
+ total = max(len(left), len(right), 1)
186
+ return {"matches": matches, "agreement_rate": matches / total, "mean_absolute_error_s": sum(errors) / len(errors) if errors else None}
187
+
188
+
189
+ def interval_iou(a: Annotation, b: Annotation) -> float:
190
+ if a.type != "interval" or b.type != "interval":
191
+ return 0.0
192
+ inter = max(0.0, min(a.end, b.end) - max(a.start, b.start)) # type: ignore[arg-type]
193
+ union = max(a.end, b.end) - min(a.start, b.start) # type: ignore[arg-type]
194
+ return inter / union if union else 0.0
@@ -0,0 +1,122 @@
1
+ """Classical open R-peak detectors for locked baseline comparison.
2
+
3
+ Implementations follow the classic literature at a research level:
4
+ - Pan & Tompkins (1985): bandpass → derivative → square → moving integrate → adaptive threshold
5
+ - Hamilton & Tompkins (1986): similar pipeline with Hamilton-style peak decision rules
6
+
7
+ These are retrospective full-record detectors (same evaluation mode as ElectroTrace
8
+ Stage-1) and are intended only for protocol-matched comparison on MIT-BIH.
9
+ """
10
+ from __future__ import annotations
11
+
12
+ import numpy as np
13
+ from scipy import signal as sps
14
+
15
+
16
+ def _validate(signal: np.ndarray, fs_hz: float) -> tuple[np.ndarray, float]:
17
+ x = np.asarray(signal, dtype=float)
18
+ fs = float(fs_hz)
19
+ if x.ndim != 1 or x.size < 32:
20
+ raise ValueError("signal must be one-dimensional with at least 32 samples")
21
+ if not np.isfinite(x).all():
22
+ raise ValueError("signal must contain only finite values")
23
+ if not np.isfinite(fs) or fs <= 0:
24
+ raise ValueError("fs_hz must be positive and finite")
25
+ return x, fs
26
+
27
+
28
+ def _bandpass(x: np.ndarray, fs_hz: float, low: float = 5.0, high: float = 15.0) -> np.ndarray:
29
+ nyq = fs_hz / 2.0
30
+ low = min(low, nyq * 0.45)
31
+ high = min(high, nyq * 0.9)
32
+ if high <= low or low <= 0:
33
+ return x - np.median(x)
34
+ sos = sps.butter(2, [low, high], btype="bandpass", fs=fs_hz, output="sos")
35
+ try:
36
+ return sps.sosfiltfilt(sos, x - np.median(x))
37
+ except ValueError:
38
+ return sps.sosfilt(sos, x - np.median(x))
39
+
40
+
41
+ def pan_tompkins_r_peaks(signal: np.ndarray, fs_hz: float) -> np.ndarray:
42
+ """Pan–Tompkins style R-peak detector (full-record retrospective)."""
43
+ x, fs = _validate(signal, fs_hz)
44
+ band = _bandpass(x, fs, 5.0, 15.0)
45
+ deriv = np.convolve(band, np.array([1, 2, 0, -2, -1], dtype=float) / 8.0, mode="same") * fs
46
+ squared = deriv ** 2
47
+ win = max(1, int(round(0.150 * fs)))
48
+ integrated = np.convolve(squared, np.ones(win, dtype=float) / win, mode="same")
49
+
50
+ init_n = min(len(integrated), int(2.0 * fs))
51
+ peak_level = float(np.max(integrated[:init_n])) if init_n else 0.0
52
+ noise_level = float(np.median(integrated[:init_n])) if init_n else 0.0
53
+ threshold = noise_level + 0.25 * (peak_level - noise_level)
54
+
55
+ min_distance = max(1, int(round(0.2 * fs)))
56
+ candidates, props = sps.find_peaks(integrated, distance=min_distance, height=threshold * 0.5)
57
+ if len(candidates) == 0:
58
+ return np.asarray([], dtype=int)
59
+
60
+ peaks: list[int] = []
61
+ last = -min_distance
62
+ for idx in candidates:
63
+ val = float(integrated[idx])
64
+ if idx - last < min_distance:
65
+ continue
66
+ if val >= threshold:
67
+ half = max(1, int(round(0.075 * fs)))
68
+ lo = max(0, int(idx) - half)
69
+ hi = min(len(band), int(idx) + half + 1)
70
+ local = band[lo:hi]
71
+ if local.size == 0:
72
+ continue
73
+ local_peak = lo + int(np.argmax(np.abs(local)))
74
+ peaks.append(local_peak)
75
+ last = idx
76
+ peak_level = 0.125 * val + 0.875 * peak_level
77
+ threshold = noise_level + 0.25 * (peak_level - noise_level)
78
+ else:
79
+ noise_level = 0.125 * val + 0.875 * noise_level
80
+ threshold = noise_level + 0.25 * (peak_level - noise_level)
81
+
82
+ return np.asarray(peaks, dtype=int)
83
+
84
+
85
+ def hamilton_r_peaks(signal: np.ndarray, fs_hz: float) -> np.ndarray:
86
+ """Hamilton–Tompkins style R-peak detector (full-record retrospective)."""
87
+ x, fs = _validate(signal, fs_hz)
88
+ band = _bandpass(x, fs, 8.0, 16.0)
89
+ deriv = np.diff(band, prepend=band[0]) * fs
90
+ squared = deriv ** 2
91
+ win = max(1, int(round(0.080 * fs)))
92
+ integrated = np.convolve(squared, np.ones(win, dtype=float) / win, mode="same")
93
+
94
+ min_distance = max(1, int(round(0.2 * fs)))
95
+ med = float(np.median(integrated))
96
+ mad = float(np.median(np.abs(integrated - med)))
97
+ scale = 1.4826 * mad if mad > 1e-12 else float(np.std(integrated) + 1e-12)
98
+ height = med + 2.5 * scale
99
+
100
+ candidates, _ = sps.find_peaks(integrated, distance=min_distance, height=height)
101
+ if len(candidates) == 0:
102
+ candidates, _ = sps.find_peaks(integrated, distance=min_distance, height=med + 1.0 * scale)
103
+
104
+ peaks: list[int] = []
105
+ half = max(1, int(round(0.06 * fs)))
106
+ for idx in candidates:
107
+ lo = max(0, int(idx) - half)
108
+ hi = min(len(band), int(idx) + half + 1)
109
+ local = band[lo:hi]
110
+ if local.size == 0:
111
+ continue
112
+ local_peak = lo + int(np.argmax(np.abs(local)))
113
+ peaks.append(local_peak)
114
+ if not peaks:
115
+ return np.asarray([], dtype=int)
116
+ return np.asarray(sorted(set(peaks)), dtype=int)
117
+
118
+
119
+ BASELINE_DETECTORS = {
120
+ "pan_tompkins": pan_tompkins_r_peaks,
121
+ "hamilton": hamilton_r_peaks,
122
+ }
electrotrace/beats.py ADDED
@@ -0,0 +1,47 @@
1
+ """Beat-level ECG segmentation from detected R peaks."""
2
+ from __future__ import annotations
3
+
4
+ from dataclasses import asdict, dataclass
5
+ import numpy as np
6
+
7
+
8
+ @dataclass(frozen=True)
9
+ class Beat:
10
+ index: int
11
+ r_index: int
12
+ r_time: float
13
+ start: float
14
+ end: float
15
+ rr_prev_s: float | None
16
+ rr_next_s: float | None
17
+ heart_rate_bpm: float | None
18
+
19
+ def to_dict(self) -> dict:
20
+ return asdict(self)
21
+
22
+
23
+ def segment_beats(time: np.ndarray, peaks: np.ndarray, pre_s: float = 0.35, post_s: float = 0.55) -> list[Beat]:
24
+ time = np.asarray(time, dtype=float)
25
+ peaks = np.asarray(peaks, dtype=int)
26
+ if time.ndim != 1 or len(time) < 2:
27
+ raise ValueError("time must be a one-dimensional array with at least two samples")
28
+ if not np.isfinite(time).all() or np.any(np.diff(time) <= 0):
29
+ raise ValueError("time must be finite and strictly increasing")
30
+ if len(peaks) == 0:
31
+ return []
32
+ if not np.isfinite(pre_s) or not np.isfinite(post_s) or pre_s <= 0 or post_s <= 0:
33
+ raise ValueError("beat windows must be positive and finite")
34
+ peaks = np.unique(peaks[(peaks >= 0) & (peaks < len(time))])
35
+ if len(peaks) == 0:
36
+ return []
37
+ r_times = time[peaks]
38
+ beats: list[Beat] = []
39
+ for i, (idx, rt) in enumerate(zip(peaks, r_times)):
40
+ prev_rr = float(rt - r_times[i - 1]) if i else None
41
+ next_rr = float(r_times[i + 1] - rt) if i + 1 < len(r_times) else None
42
+ rr = prev_rr if prev_rr is not None and prev_rr > 0 else next_rr
43
+ hr = 60.0 / rr if rr is not None and rr > 0 else None
44
+ start = max(float(time[0]), float(rt - pre_s))
45
+ end = min(float(time[-1]), float(rt + post_s))
46
+ beats.append(Beat(i, int(idx), float(rt), start, end, prev_rr, next_rr, hr))
47
+ return beats
@@ -0,0 +1,154 @@
1
+ """Leakage-safe model benchmarking with subject-level, stratified splits."""
2
+ from __future__ import annotations
3
+
4
+ from dataclasses import asdict, dataclass
5
+
6
+ import numpy as np
7
+ from scipy.stats import t as student_t
8
+ from sklearn.ensemble import RandomForestClassifier
9
+ from sklearn.linear_model import LogisticRegression
10
+ from sklearn.metrics import accuracy_score, balanced_accuracy_score, f1_score, roc_auc_score, confusion_matrix
11
+ from sklearn.model_selection import StratifiedGroupKFold
12
+ from sklearn.pipeline import make_pipeline
13
+ from sklearn.preprocessing import StandardScaler
14
+
15
+
16
+ @dataclass
17
+ class FoldMetrics:
18
+ fold: int
19
+ n_train: int
20
+ n_test: int
21
+ accuracy: float
22
+ balanced_accuracy: float
23
+ macro_f1: float
24
+ weighted_f1: float
25
+ roc_auc: float | None
26
+
27
+ def to_dict(self) -> dict:
28
+ return asdict(self)
29
+
30
+
31
+ def _model_classes(model) -> np.ndarray:
32
+ classes = getattr(model, "classes_", None)
33
+ if classes is not None:
34
+ return np.asarray(classes)
35
+ steps = getattr(model, "steps", None)
36
+ if steps:
37
+ estimator = steps[-1][1]
38
+ classes = getattr(estimator, "classes_", None)
39
+ if classes is not None:
40
+ return np.asarray(classes)
41
+ raise ValueError("trained model does not expose class labels")
42
+
43
+
44
+ def _auc(y_true: np.ndarray, proba: np.ndarray, classes: np.ndarray) -> float | None:
45
+ try:
46
+ if len(classes) == 2:
47
+ return float(roc_auc_score(y_true, proba[:, 1]))
48
+ return float(roc_auc_score(y_true, proba, multi_class="ovr", labels=classes))
49
+ except ValueError:
50
+ return None
51
+
52
+
53
+ def _summary(values: list[float | None]) -> dict[str, float | None]:
54
+ """Describe cross-validation fold variability.
55
+
56
+ The t-based interval is explicitly descriptive across CV folds. Because CV
57
+ training sets overlap, it must not be interpreted as an independent-sample
58
+ confidence interval for future-subject performance.
59
+ """
60
+ x = np.asarray([v for v in values if v is not None and np.isfinite(v)], dtype=float)
61
+ if x.size == 0:
62
+ return {"n": 0, "mean": None, "std": None, "ci95_low": None, "ci95_high": None}
63
+ mean = float(np.mean(x))
64
+ if x.size == 1:
65
+ return {"n": 1, "mean": mean, "std": None, "ci95_low": None, "ci95_high": None}
66
+ sd = float(np.std(x, ddof=1))
67
+ critical = float(student_t.ppf(0.975, df=x.size - 1))
68
+ margin = critical * sd / np.sqrt(x.size)
69
+ return {"n": int(x.size), "mean": mean, "std": sd, "ci95_low": mean - margin, "ci95_high": mean + margin}
70
+
71
+
72
+ def benchmark_models(X: np.ndarray, y: np.ndarray, groups: np.ndarray, folds: int = 5, seed: int = 42) -> dict:
73
+ """Benchmark baseline classifiers using leakage-safe group-stratified folds.
74
+
75
+ ``groups`` are the experimental units for splitting. Samples from one group
76
+ never occur in both train and test within a fold.
77
+ """
78
+ X = np.asarray(X, dtype=float)
79
+ y = np.asarray(y)
80
+ groups = np.asarray(groups)
81
+ if X.ndim != 2:
82
+ raise ValueError("X must be a two-dimensional feature matrix")
83
+ if not np.isfinite(X).all():
84
+ raise ValueError("X contains NaN or infinite values")
85
+ if len(X) != len(y) or len(y) != len(groups):
86
+ raise ValueError("X, y, and groups must have equal length")
87
+ labels = np.unique(y)
88
+ unique_groups = np.unique(groups)
89
+ if len(labels) < 2:
90
+ raise ValueError("Need at least two outcome classes for classification benchmarking")
91
+ if len(unique_groups) < 2:
92
+ raise ValueError("Need at least two subjects/groups for leakage-safe benchmarking")
93
+ requested_folds = int(folds)
94
+ if requested_folds < 2:
95
+ raise ValueError("folds must be at least 2")
96
+ groups_per_class = {str(cls): len(np.unique(groups[y == cls])) for cls in labels}
97
+ max_folds = min(len(unique_groups), min(groups_per_class.values()))
98
+ n_splits = min(requested_folds, max_folds)
99
+ if n_splits < 2:
100
+ raise ValueError(f"Need at least two distinct subjects in every class; observed {groups_per_class}")
101
+
102
+ models = {
103
+ "logistic_regression": make_pipeline(StandardScaler(), LogisticRegression(max_iter=2000, class_weight="balanced", random_state=seed)),
104
+ "random_forest": RandomForestClassifier(n_estimators=300, class_weight="balanced", random_state=seed, n_jobs=-1),
105
+ }
106
+ splitter = StratifiedGroupKFold(n_splits=n_splits, shuffle=True, random_state=seed)
107
+ result: dict[str, object] = {
108
+ "n_samples": int(len(y)),
109
+ "n_subjects": int(len(unique_groups)),
110
+ "experimental_unit": "group",
111
+ "folds": n_splits,
112
+ "splitter": "StratifiedGroupKFold",
113
+ "seed": int(seed),
114
+ "groups_per_class": groups_per_class,
115
+ "models": {},
116
+ }
117
+ metric_names = ("accuracy", "balanced_accuracy", "macro_f1", "weighted_f1", "roc_auc")
118
+ for name, model in models.items():
119
+ metrics: list[FoldMetrics] = []
120
+ confusion = None
121
+ test_groups_per_fold: list[int] = []
122
+ for fold, (train_idx, test_idx) in enumerate(splitter.split(X, y, groups), start=1):
123
+ if len(np.unique(y[train_idx])) < 2:
124
+ raise ValueError(f"Fold {fold} training set contains fewer than two outcome classes")
125
+ if set(groups[train_idx]).intersection(set(groups[test_idx])):
126
+ raise ValueError(f"Fold {fold} leaks experimental units between train and test")
127
+ model.fit(X[train_idx], y[train_idx])
128
+ pred = model.predict(X[test_idx])
129
+ proba = model.predict_proba(X[test_idx])
130
+ classes = _model_classes(model)
131
+ metrics.append(FoldMetrics(
132
+ fold=fold,
133
+ n_train=len(train_idx), n_test=len(test_idx),
134
+ accuracy=float(accuracy_score(y[test_idx], pred)),
135
+ balanced_accuracy=float(balanced_accuracy_score(y[test_idx], pred)),
136
+ macro_f1=float(f1_score(y[test_idx], pred, average="macro", zero_division=0)),
137
+ weighted_f1=float(f1_score(y[test_idx], pred, average="weighted", zero_division=0)),
138
+ roc_auc=_auc(y[test_idx], proba, classes),
139
+ ))
140
+ test_groups_per_fold.append(int(len(np.unique(groups[test_idx]))))
141
+ cm = confusion_matrix(y[test_idx], pred, labels=labels)
142
+ confusion = cm if confusion is None else confusion + cm
143
+ fold_values = {k: [getattr(m, k) for m in metrics] for k in metric_names}
144
+ summaries = {k: _summary(v) for k, v in fold_values.items()}
145
+ result["models"][name] = {
146
+ "folds": [m.to_dict() for m in metrics],
147
+ "summary": summaries,
148
+ "summary_interval_interpretation": "descriptive_t_interval_across_CV_folds; not an independent-subject generalization CI",
149
+ "mean": {k: s["mean"] for k, s in summaries.items()},
150
+ "confusion_matrix": confusion.tolist() if confusion is not None else [],
151
+ "labels": [str(x) for x in labels],
152
+ "test_groups_per_fold": test_groups_per_fold,
153
+ }
154
+ return result