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.
- electrotrace/__init__.py +28 -0
- electrotrace/__main__.py +3 -0
- electrotrace/annotations.py +194 -0
- electrotrace/baseline_detectors.py +122 -0
- electrotrace/beats.py +47 -0
- electrotrace/benchmark.py +154 -0
- electrotrace/candidate_suppressor.py +441 -0
- electrotrace/cli.py +630 -0
- electrotrace/detectors.py +152 -0
- electrotrace/formats.py +147 -0
- electrotrace/fp_analysis.py +358 -0
- electrotrace/hearttwin_adapter.py +90 -0
- electrotrace/io.py +163 -0
- electrotrace/lead_quality.py +138 -0
- electrotrace/lead_selection.py +169 -0
- electrotrace/metadata.py +111 -0
- electrotrace/ml.py +178 -0
- electrotrace/phenotype.py +84 -0
- electrotrace/phenotype_validation.py +133 -0
- electrotrace/polarity_v2.py +201 -0
- electrotrace/project.py +43 -0
- electrotrace/project_store.py +147 -0
- electrotrace/provenance.py +155 -0
- electrotrace/qrs_delineation.py +139 -0
- electrotrace/qtdb_detector_adapter.py +7 -0
- electrotrace/research_validation.py +146 -0
- electrotrace/scale_estimation.py +172 -0
- electrotrace/security.py +57 -0
- electrotrace/server_app.py +375 -0
- electrotrace/signal.py +84 -0
- electrotrace/statistics.py +98 -0
- electrotrace/threshold_selection.py +40 -0
- electrotrace/validation.py +241 -0
- electrotrace/validation_detectors.py +570 -0
- electrotrace/wfdb_records.py +375 -0
- electrotrace/window.py +117 -0
- electrotrace-1.9.0.dist-info/METADATA +187 -0
- electrotrace-1.9.0.dist-info/RECORD +42 -0
- electrotrace-1.9.0.dist-info/WHEEL +5 -0
- electrotrace-1.9.0.dist-info/entry_points.txt +3 -0
- electrotrace-1.9.0.dist-info/licenses/LICENSE +20 -0
- electrotrace-1.9.0.dist-info/top_level.txt +1 -0
electrotrace/__init__.py
ADDED
|
@@ -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
|
+
]
|
electrotrace/__main__.py
ADDED
|
@@ -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
|