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
|
@@ -0,0 +1,441 @@
|
|
|
1
|
+
"""Second-stage false-positive suppression for ECG R-peak candidates."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from dataclasses import asdict, dataclass, replace
|
|
5
|
+
from pathlib import Path
|
|
6
|
+
import hashlib
|
|
7
|
+
import json
|
|
8
|
+
import pickle
|
|
9
|
+
import warnings
|
|
10
|
+
from typing import Sequence
|
|
11
|
+
|
|
12
|
+
import numpy as np
|
|
13
|
+
from scipy import signal as sps
|
|
14
|
+
from sklearn.ensemble import RandomForestClassifier
|
|
15
|
+
from sklearn.model_selection import StratifiedShuffleSplit
|
|
16
|
+
|
|
17
|
+
from .scale_estimation import DEFAULT_SCALE_METHOD, estimate_scale
|
|
18
|
+
|
|
19
|
+
SUPPRESSOR_VERSION = "rf-candidate-suppressor-v4"
|
|
20
|
+
FEATURE_SCHEMA_VERSION = "candidate-features-v4" # bumped: feature normalization changed (see below)
|
|
21
|
+
MODEL_FORMAT_VERSION = "electrotrace-model-v2-skops"
|
|
22
|
+
DEFAULT_TRUSTED_SKOPS_TYPES = (
|
|
23
|
+
"numpy.core.multiarray._reconstruct",
|
|
24
|
+
"numpy.core.multiarray.scalar",
|
|
25
|
+
"numpy.dtype",
|
|
26
|
+
"numpy.ndarray",
|
|
27
|
+
"sklearn.ensemble._forest.RandomForestClassifier",
|
|
28
|
+
"sklearn.tree._classes.DecisionTreeClassifier",
|
|
29
|
+
"sklearn.tree._tree.Tree",
|
|
30
|
+
)
|
|
31
|
+
DEFAULT_TARGET_RECALL = 0.995
|
|
32
|
+
DEFAULT_TOLERANCE_S = 0.075
|
|
33
|
+
DEFAULT_CALIBRATION_FRACTION = 0.20
|
|
34
|
+
FEATURE_CHUNK_SIZE = 4096
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
@dataclass(frozen=True)
|
|
38
|
+
class SuppressorMetadata:
|
|
39
|
+
model_version: str
|
|
40
|
+
feature_schema_version: str
|
|
41
|
+
target_recall: float
|
|
42
|
+
threshold: float
|
|
43
|
+
n_training_candidates: int
|
|
44
|
+
n_positive_candidates: int
|
|
45
|
+
n_negative_candidates: int
|
|
46
|
+
random_seed: int
|
|
47
|
+
n_estimators: int
|
|
48
|
+
calibration_fraction: float = DEFAULT_CALIBRATION_FRACTION
|
|
49
|
+
calibration_candidates: int = 0
|
|
50
|
+
calibration_method: str = "held_out_stratified"
|
|
51
|
+
sklearn_version: str = ""
|
|
52
|
+
|
|
53
|
+
def to_dict(self) -> dict:
|
|
54
|
+
return asdict(self)
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _validate_signal(signal: np.ndarray, fs_hz: float) -> tuple[np.ndarray, float]:
|
|
58
|
+
signal = np.asarray(signal, dtype=float)
|
|
59
|
+
fs_hz = float(fs_hz)
|
|
60
|
+
if signal.ndim != 1 or signal.size < 16:
|
|
61
|
+
raise ValueError("signal must be one-dimensional with at least 16 samples")
|
|
62
|
+
if not np.isfinite(signal).all():
|
|
63
|
+
raise ValueError("signal must contain only finite values")
|
|
64
|
+
if not np.isfinite(fs_hz) or fs_hz <= 0:
|
|
65
|
+
raise ValueError("fs_hz must be positive and finite")
|
|
66
|
+
return signal, fs_hz
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _bandpasses(zsignal: np.ndarray, fs_hz: float) -> list[np.ndarray]:
|
|
70
|
+
nyq = fs_hz / 2.0
|
|
71
|
+
specs = [(0.5, min(5.0, nyq * 0.9)), (5.0, min(15.0, nyq * 0.9)),
|
|
72
|
+
(15.0, min(40.0, nyq * 0.9)), (40.0, min(100.0, nyq * 0.9))]
|
|
73
|
+
out: list[np.ndarray] = []
|
|
74
|
+
for low, high in specs:
|
|
75
|
+
if high <= low or high <= 0:
|
|
76
|
+
out.append(np.zeros_like(zsignal))
|
|
77
|
+
continue
|
|
78
|
+
sos = sps.butter(3, [low, high], btype="bandpass", fs=fs_hz, output="sos")
|
|
79
|
+
try:
|
|
80
|
+
out.append(sps.sosfiltfilt(sos, zsignal))
|
|
81
|
+
except ValueError:
|
|
82
|
+
out.append(np.zeros_like(zsignal))
|
|
83
|
+
return out
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _candidate_features(
|
|
87
|
+
signal: np.ndarray,
|
|
88
|
+
fs_hz: float,
|
|
89
|
+
candidate_indices: Sequence[int],
|
|
90
|
+
prominences: Sequence[float] | None = None,
|
|
91
|
+
window_s: float = 0.25,
|
|
92
|
+
*,
|
|
93
|
+
scale_method: str = DEFAULT_SCALE_METHOD,
|
|
94
|
+
) -> tuple[np.ndarray, list[str]]:
|
|
95
|
+
"""Extract candidate features in bounded vectorized chunks.
|
|
96
|
+
|
|
97
|
+
scale_method controls how `global_scale` (used to normalize both the
|
|
98
|
+
z-signal and prominences below) is estimated. Originally a plain global
|
|
99
|
+
std, which has the same fragile-global-statistic problem Stage-1 had:
|
|
100
|
+
for a record with severe multi-minute DC drift (e.g. INCART I03, ~8mV
|
|
101
|
+
drift vs ~0.3-0.5mV true QRS amplitude), global std is dominated by the
|
|
102
|
+
drift, so true-beat prominences normalize to near-zero and look like
|
|
103
|
+
noise to the RF -- even when Stage-1 generates plenty of correct
|
|
104
|
+
candidates. Defaulting to the same windowed_std estimator used in
|
|
105
|
+
Stage-1 fixes this (see scale_estimation.py for the full history).
|
|
106
|
+
Changing this default changes the feature space the RF is trained on,
|
|
107
|
+
hence the FEATURE_SCHEMA_VERSION bump above -- requires a full retrain,
|
|
108
|
+
not just reloading an old model.
|
|
109
|
+
"""
|
|
110
|
+
signal, fs_hz = _validate_signal(signal, fs_hz)
|
|
111
|
+
candidates = np.asarray(candidate_indices, dtype=int)
|
|
112
|
+
if candidates.ndim != 1:
|
|
113
|
+
raise ValueError("candidate_indices must be one-dimensional")
|
|
114
|
+
if np.any(candidates < 0) or np.any(candidates >= signal.size):
|
|
115
|
+
raise ValueError("candidate_indices contain out-of-range values")
|
|
116
|
+
if not np.isfinite(window_s) or window_s <= 0:
|
|
117
|
+
raise ValueError("window_s must be positive and finite")
|
|
118
|
+
|
|
119
|
+
base = signal - np.median(signal)
|
|
120
|
+
global_scale = estimate_scale(base, fs_hz, method=scale_method) or 1.0
|
|
121
|
+
zsignal = base / global_scale
|
|
122
|
+
prominences = np.zeros(len(candidates), dtype=float) if prominences is None else np.asarray(prominences, dtype=float)
|
|
123
|
+
if prominences.ndim != 1 or len(prominences) != len(candidates):
|
|
124
|
+
raise ValueError("prominences must match candidate_indices")
|
|
125
|
+
if not np.isfinite(prominences).all():
|
|
126
|
+
raise ValueError("prominences must be finite")
|
|
127
|
+
|
|
128
|
+
rr = np.diff(candidates) / fs_hz if len(candidates) > 1 else np.asarray([], dtype=float)
|
|
129
|
+
rr_median = float(np.median(rr)) if rr.size else 1.0
|
|
130
|
+
band_signals = _bandpasses(zsignal, fs_hz)
|
|
131
|
+
|
|
132
|
+
names = [
|
|
133
|
+
"amplitude_z", "prominence_z", "width_s", "max_abs_slope", "mean_abs_slope",
|
|
134
|
+
"local_rms", "crest_factor", "low_band_fraction", "qrs_band_fraction",
|
|
135
|
+
"high_band_fraction", "very_high_band_fraction", "left_right_energy_ratio",
|
|
136
|
+
"left_right_amplitude_ratio", "rr_prev_s", "rr_next_s", "rr_prev_ratio",
|
|
137
|
+
"rr_next_ratio",
|
|
138
|
+
] + [f"shape_{i:02d}" for i in range(64)]
|
|
139
|
+
|
|
140
|
+
radius = max(16, int(round(fs_hz * window_s)))
|
|
141
|
+
half_qrs = max(4, int(round(fs_hz * 0.08)))
|
|
142
|
+
offsets = np.arange(-radius, radius + 1, dtype=int)
|
|
143
|
+
qslice = slice(radius - half_qrs, radius + half_qrs + 1)
|
|
144
|
+
leftslice = slice(radius - half_qrs, radius)
|
|
145
|
+
rightslice = slice(radius + 1, radius + half_qrs + 1)
|
|
146
|
+
shape_pos = np.linspace(0.0, len(offsets) - 1.0, 64)
|
|
147
|
+
shape_lo = np.floor(shape_pos).astype(int)
|
|
148
|
+
shape_hi = np.minimum(shape_lo + 1, len(offsets) - 1)
|
|
149
|
+
shape_alpha = (shape_pos - shape_lo).astype(np.float32)
|
|
150
|
+
|
|
151
|
+
rows: list[np.ndarray] = []
|
|
152
|
+
n = len(candidates)
|
|
153
|
+
for start in range(0, n, FEATURE_CHUNK_SIZE):
|
|
154
|
+
stop = min(n, start + FEATURE_CHUNK_SIZE)
|
|
155
|
+
cand = candidates[start:stop]
|
|
156
|
+
idx = np.clip(cand[:, None] + offsets[None, :], 0, signal.size - 1)
|
|
157
|
+
x = zsignal[idx]
|
|
158
|
+
shape = x[:, shape_lo] * (1.0 - shape_alpha[None, :]) + x[:, shape_hi] * shape_alpha[None, :]
|
|
159
|
+
derivative = np.diff(x, axis=1) * fs_hz
|
|
160
|
+
qrs = x[:, qslice]
|
|
161
|
+
local_rms = np.sqrt(np.mean(qrs * qrs, axis=1))
|
|
162
|
+
peak_abs = np.max(np.abs(qrs), axis=1)
|
|
163
|
+
mean_abs = np.mean(np.abs(qrs), axis=1)
|
|
164
|
+
crest = peak_abs / np.maximum(mean_abs, 1e-6)
|
|
165
|
+
width = np.count_nonzero(np.abs(qrs) >= (0.5 * peak_abs)[:, None], axis=1) / fs_hz
|
|
166
|
+
|
|
167
|
+
left = x[:, leftslice]
|
|
168
|
+
right = x[:, rightslice]
|
|
169
|
+
left_energy = np.mean(left * left, axis=1)
|
|
170
|
+
right_energy = np.mean(right * right, axis=1)
|
|
171
|
+
left_amp = np.max(np.abs(left), axis=1)
|
|
172
|
+
right_amp = np.max(np.abs(right), axis=1)
|
|
173
|
+
|
|
174
|
+
fractions = []
|
|
175
|
+
bandq = [band[idx[:, qslice]] for band in band_signals]
|
|
176
|
+
for b in bandq:
|
|
177
|
+
fractions.append(np.mean(b * b, axis=1))
|
|
178
|
+
total_band = np.sum(fractions, axis=0)
|
|
179
|
+
total_band = np.maximum(total_band, 1e-8)
|
|
180
|
+
bands = [v / total_band for v in fractions]
|
|
181
|
+
|
|
182
|
+
prev_rr = np.empty(stop - start, dtype=float)
|
|
183
|
+
next_rr = np.empty(stop - start, dtype=float)
|
|
184
|
+
if stop - start:
|
|
185
|
+
prev_rr[0] = rr_median
|
|
186
|
+
if stop - start > 1:
|
|
187
|
+
prev_rr[1:] = np.diff(cand) / fs_hz
|
|
188
|
+
next_rr[:-1] = np.diff(cand) / fs_hz
|
|
189
|
+
next_rr[-1] = rr_median
|
|
190
|
+
prev_ratio = prev_rr / max(rr_median, 1e-6)
|
|
191
|
+
next_ratio = next_rr / max(rr_median, 1e-6)
|
|
192
|
+
|
|
193
|
+
chunk = np.column_stack([
|
|
194
|
+
zsignal[cand], prominences[start:stop] / global_scale, width,
|
|
195
|
+
np.max(np.abs(derivative), axis=1), np.mean(np.abs(derivative), axis=1),
|
|
196
|
+
local_rms, crest, *bands,
|
|
197
|
+
left_energy / np.maximum(right_energy, 1e-6),
|
|
198
|
+
left_amp / np.maximum(right_amp, 1e-6),
|
|
199
|
+
prev_rr, next_rr, prev_ratio, next_ratio,
|
|
200
|
+
shape,
|
|
201
|
+
]).astype(np.float32)
|
|
202
|
+
rows.append(chunk)
|
|
203
|
+
|
|
204
|
+
if not rows:
|
|
205
|
+
return np.empty((0, len(names)), dtype=np.float32), names
|
|
206
|
+
out = np.vstack(rows)
|
|
207
|
+
if not np.isfinite(out).all():
|
|
208
|
+
raise ValueError("candidate feature extraction produced NaN or infinite values")
|
|
209
|
+
return out, names
|
|
210
|
+
|
|
211
|
+
|
|
212
|
+
def label_candidates(candidate_indices: Sequence[int], reference_indices: Sequence[int], fs_hz: float,
|
|
213
|
+
tolerance_s: float = DEFAULT_TOLERANCE_S) -> np.ndarray:
|
|
214
|
+
fs_hz = float(fs_hz); tolerance_s = float(tolerance_s)
|
|
215
|
+
candidates = np.asarray(candidate_indices, dtype=int); references = np.asarray(reference_indices, dtype=int)
|
|
216
|
+
if fs_hz <= 0 or not np.isfinite(fs_hz): raise ValueError("fs_hz must be positive and finite")
|
|
217
|
+
if tolerance_s <= 0 or not np.isfinite(tolerance_s): raise ValueError("tolerance_s must be positive and finite")
|
|
218
|
+
if candidates.ndim != 1 or references.ndim != 1: raise ValueError("candidate_indices and reference_indices must be one-dimensional")
|
|
219
|
+
if not np.all(candidates[1:] >= candidates[:-1]) or not np.all(references[1:] >= references[:-1]):
|
|
220
|
+
raise ValueError("candidate_indices and reference_indices must be sorted")
|
|
221
|
+
labels = np.zeros(len(candidates), dtype=np.int8); used_reference = np.zeros(len(references), dtype=bool)
|
|
222
|
+
tolerance_samples = tolerance_s * fs_hz
|
|
223
|
+
for i, candidate in enumerate(candidates):
|
|
224
|
+
pos = int(np.searchsorted(references, candidate)); possible = []
|
|
225
|
+
if pos < len(references): possible.append(pos)
|
|
226
|
+
if pos > 0: possible.append(pos - 1)
|
|
227
|
+
if possible:
|
|
228
|
+
best = min(possible, key=lambda j: abs(int(references[j]) - int(candidate)))
|
|
229
|
+
if not used_reference[best] and abs(int(references[best]) - int(candidate)) <= tolerance_samples:
|
|
230
|
+
labels[i] = 1; used_reference[best] = True
|
|
231
|
+
return labels
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
def select_threshold_for_recall(y_true: Sequence[int], probabilities: Sequence[float], target_recall: float = DEFAULT_TARGET_RECALL) -> float:
|
|
235
|
+
y_true = np.asarray(y_true, dtype=int); probabilities = np.asarray(probabilities, dtype=float); target_recall = float(target_recall)
|
|
236
|
+
if y_true.shape != probabilities.shape or y_true.ndim != 1: raise ValueError("y_true and probabilities must be one-dimensional and equal length")
|
|
237
|
+
if not (0 < target_recall <= 1) or not np.isfinite(target_recall): raise ValueError("target_recall must be in (0, 1]")
|
|
238
|
+
if not np.isfinite(probabilities).all(): raise ValueError("probabilities must be finite")
|
|
239
|
+
positives = int(y_true.sum())
|
|
240
|
+
if positives == 0: raise ValueError("at least one positive candidate is required")
|
|
241
|
+
order = np.argsort(-probabilities); recall = np.cumsum(y_true[order]) / positives
|
|
242
|
+
valid = np.flatnonzero(recall >= target_recall)
|
|
243
|
+
return float(probabilities[order[valid[0]]]) if valid.size else 0.0
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
class CandidateSuppressor:
|
|
247
|
+
def __init__(self, model: RandomForestClassifier | None = None, metadata: SuppressorMetadata | None = None):
|
|
248
|
+
self.model = model; self.metadata = metadata; self.feature_names: list[str] | None = None
|
|
249
|
+
|
|
250
|
+
@property
|
|
251
|
+
def fitted(self) -> bool:
|
|
252
|
+
return self.model is not None and self.metadata is not None
|
|
253
|
+
|
|
254
|
+
def fit(self, feature_matrix: np.ndarray, labels: Sequence[int], *, target_recall: float = DEFAULT_TARGET_RECALL,
|
|
255
|
+
random_seed: int = 42, n_estimators: int = 150, calibration_fraction: float = DEFAULT_CALIBRATION_FRACTION) -> "CandidateSuppressor":
|
|
256
|
+
X = np.asarray(feature_matrix, dtype=float); y = np.asarray(labels, dtype=int)
|
|
257
|
+
if X.ndim != 2 or X.shape[0] == 0: raise ValueError("feature_matrix must be a non-empty two-dimensional array")
|
|
258
|
+
if len(y) != len(X): raise ValueError("feature_matrix and labels must have equal length")
|
|
259
|
+
if not np.isfinite(X).all(): raise ValueError("feature_matrix contains NaN or infinite values")
|
|
260
|
+
if set(np.unique(y)) - {0, 1} or len(np.unique(y)) < 2: raise ValueError("labels must contain both negative and positive candidates")
|
|
261
|
+
calibration_fraction = float(calibration_fraction)
|
|
262
|
+
if not 0 <= calibration_fraction < 0.5 or not np.isfinite(calibration_fraction):
|
|
263
|
+
raise ValueError("calibration_fraction must be finite and in [0, 0.5)")
|
|
264
|
+
unique, counts = np.unique(y, return_counts=True)
|
|
265
|
+
if calibration_fraction > 0 and (len(y) < 10 or np.min(counts) < 2):
|
|
266
|
+
raise ValueError("held-out threshold calibration requires at least 10 candidates and two examples per class; set calibration_fraction=0 to disable calibration")
|
|
267
|
+
|
|
268
|
+
calibration_candidates = 0
|
|
269
|
+
calibration_method = "disabled"
|
|
270
|
+
if calibration_fraction > 0:
|
|
271
|
+
splitter = StratifiedShuffleSplit(n_splits=1, test_size=calibration_fraction, random_state=int(random_seed))
|
|
272
|
+
model_idx, calibration_idx = next(splitter.split(X, y))
|
|
273
|
+
calibration_model = RandomForestClassifier(
|
|
274
|
+
n_estimators=int(n_estimators), class_weight="balanced_subsample", min_samples_leaf=2,
|
|
275
|
+
max_depth=14, random_state=int(random_seed), n_jobs=-1,
|
|
276
|
+
)
|
|
277
|
+
calibration_model.fit(X[model_idx], y[model_idx])
|
|
278
|
+
calibration_probabilities = calibration_model.predict_proba(X[calibration_idx])[:, 1]
|
|
279
|
+
threshold = select_threshold_for_recall(y[calibration_idx], calibration_probabilities, target_recall=target_recall)
|
|
280
|
+
calibration_candidates = int(len(calibration_idx))
|
|
281
|
+
calibration_method = "held_out_stratified"
|
|
282
|
+
else:
|
|
283
|
+
final_for_threshold = RandomForestClassifier(
|
|
284
|
+
n_estimators=int(n_estimators), class_weight="balanced_subsample", min_samples_leaf=2,
|
|
285
|
+
max_depth=14, random_state=int(random_seed), n_jobs=-1,
|
|
286
|
+
)
|
|
287
|
+
final_for_threshold.fit(X, y)
|
|
288
|
+
threshold = select_threshold_for_recall(y, final_for_threshold.predict_proba(X)[:, 1], target_recall=target_recall)
|
|
289
|
+
calibration_method = "training_resubstitution"
|
|
290
|
+
|
|
291
|
+
model = RandomForestClassifier(n_estimators=int(n_estimators), class_weight="balanced_subsample", min_samples_leaf=2,
|
|
292
|
+
max_depth=14, random_state=int(random_seed), n_jobs=-1)
|
|
293
|
+
model.fit(X, y)
|
|
294
|
+
self.model = model
|
|
295
|
+
import sklearn
|
|
296
|
+
self.metadata = SuppressorMetadata(
|
|
297
|
+
SUPPRESSOR_VERSION,
|
|
298
|
+
FEATURE_SCHEMA_VERSION,
|
|
299
|
+
float(target_recall),
|
|
300
|
+
float(threshold),
|
|
301
|
+
int(len(y)),
|
|
302
|
+
int(y.sum()),
|
|
303
|
+
int((y == 0).sum()),
|
|
304
|
+
int(random_seed),
|
|
305
|
+
int(n_estimators),
|
|
306
|
+
calibration_fraction=float(calibration_fraction),
|
|
307
|
+
calibration_candidates=calibration_candidates,
|
|
308
|
+
calibration_method=calibration_method,
|
|
309
|
+
sklearn_version=str(sklearn.__version__),
|
|
310
|
+
)
|
|
311
|
+
return self
|
|
312
|
+
|
|
313
|
+
def predict_proba(self, features: np.ndarray) -> np.ndarray:
|
|
314
|
+
if not self.fitted: raise ValueError("candidate suppressor is not fitted")
|
|
315
|
+
X = np.asarray(features, dtype=float)
|
|
316
|
+
if X.ndim != 2 or not np.isfinite(X).all(): raise ValueError("features must be a finite two-dimensional array")
|
|
317
|
+
return self.model.predict_proba(X)[:, 1]
|
|
318
|
+
|
|
319
|
+
def filter_candidates(self, candidate_indices: Sequence[int], features: np.ndarray, threshold: float | None = None) -> tuple[np.ndarray, np.ndarray]:
|
|
320
|
+
probabilities = self.predict_proba(features)
|
|
321
|
+
if len(candidate_indices) != len(probabilities): raise ValueError("candidate_indices and features must have equal length")
|
|
322
|
+
selected_threshold = float(self.metadata.threshold if threshold is None else threshold)
|
|
323
|
+
if not 0 <= selected_threshold <= 1: raise ValueError("threshold must be in [0, 1]")
|
|
324
|
+
candidates = np.asarray(candidate_indices, dtype=int); mask = probabilities >= selected_threshold
|
|
325
|
+
return candidates[mask], probabilities[mask]
|
|
326
|
+
|
|
327
|
+
def save(self, path: str | Path) -> None:
|
|
328
|
+
"""Persist the fitted sklearn model using skops plus a JSON sidecar."""
|
|
329
|
+
if not self.fitted:
|
|
330
|
+
raise ValueError("cannot save an unfitted candidate suppressor")
|
|
331
|
+
path = Path(path)
|
|
332
|
+
if path.suffix.lower() != ".skops":
|
|
333
|
+
raise ValueError(
|
|
334
|
+
"model paths must end in .skops; legacy pickle is read-only compatibility format"
|
|
335
|
+
)
|
|
336
|
+
try:
|
|
337
|
+
from skops import io as skops_io
|
|
338
|
+
except ImportError as exc:
|
|
339
|
+
raise RuntimeError(
|
|
340
|
+
"skops is required for safe model persistence; install with pip install skops"
|
|
341
|
+
) from exc
|
|
342
|
+
skops_io.dump(self.model, path)
|
|
343
|
+
model_sha256 = hashlib.sha256(path.read_bytes()).hexdigest()
|
|
344
|
+
metadata_path = path.with_suffix(path.suffix + ".json")
|
|
345
|
+
payload = {
|
|
346
|
+
"format_version": MODEL_FORMAT_VERSION,
|
|
347
|
+
"metadata": self.metadata.to_dict(),
|
|
348
|
+
"feature_names": self.feature_names or [],
|
|
349
|
+
"model_sha256": model_sha256,
|
|
350
|
+
}
|
|
351
|
+
metadata_path.write_text(
|
|
352
|
+
json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8"
|
|
353
|
+
)
|
|
354
|
+
|
|
355
|
+
@classmethod
|
|
356
|
+
def load(
|
|
357
|
+
cls,
|
|
358
|
+
path: str | Path,
|
|
359
|
+
*,
|
|
360
|
+
allow_pickle: bool = False,
|
|
361
|
+
trusted_types: Sequence[str] = (),
|
|
362
|
+
) -> "CandidateSuppressor":
|
|
363
|
+
"""Load a model with safe-by-default serialization."""
|
|
364
|
+
import sklearn
|
|
365
|
+
|
|
366
|
+
path = Path(path)
|
|
367
|
+
if path.suffix.lower() == ".pkl":
|
|
368
|
+
if not allow_pickle:
|
|
369
|
+
raise ValueError(
|
|
370
|
+
"Refusing legacy pickle model. Convert it to .skops first; "
|
|
371
|
+
"use allow_pickle=True only for a trusted local migration."
|
|
372
|
+
)
|
|
373
|
+
warnings.warn(
|
|
374
|
+
"Loading legacy pickle; use .skops for safe persistence.",
|
|
375
|
+
RuntimeWarning,
|
|
376
|
+
stacklevel=2,
|
|
377
|
+
)
|
|
378
|
+
payload = pickle.loads(path.read_bytes())
|
|
379
|
+
metadata = payload["metadata"]
|
|
380
|
+
if not hasattr(metadata, "sklearn_version"):
|
|
381
|
+
metadata = replace(metadata, sklearn_version=str(sklearn.__version__))
|
|
382
|
+
obj = cls(model=payload["model"], metadata=metadata)
|
|
383
|
+
obj.feature_names = payload.get("feature_names")
|
|
384
|
+
return obj
|
|
385
|
+
|
|
386
|
+
if path.suffix.lower() != ".skops":
|
|
387
|
+
raise ValueError("unsupported model format; expected .skops or legacy .pkl")
|
|
388
|
+
|
|
389
|
+
metadata_path = path.with_suffix(path.suffix + ".json")
|
|
390
|
+
if not metadata_path.exists():
|
|
391
|
+
raise ValueError(f"missing model metadata sidecar: {metadata_path.name}")
|
|
392
|
+
envelope = json.loads(metadata_path.read_text(encoding="utf-8"))
|
|
393
|
+
if envelope.get("format_version") != MODEL_FORMAT_VERSION:
|
|
394
|
+
raise ValueError(
|
|
395
|
+
f"unsupported model format: {envelope.get('format_version')!r}"
|
|
396
|
+
)
|
|
397
|
+
metadata = SuppressorMetadata(**envelope["metadata"])
|
|
398
|
+
expected_sha = envelope.get("model_sha256")
|
|
399
|
+
if expected_sha:
|
|
400
|
+
actual_sha = hashlib.sha256(path.read_bytes()).hexdigest()
|
|
401
|
+
if actual_sha != expected_sha:
|
|
402
|
+
raise ValueError(
|
|
403
|
+
f"model SHA-256 mismatch for {path.name}; expected {expected_sha}, got {actual_sha}"
|
|
404
|
+
)
|
|
405
|
+
|
|
406
|
+
runtime_major = str(sklearn.__version__).split(".", 1)[0]
|
|
407
|
+
stored_major = (
|
|
408
|
+
metadata.sklearn_version.split(".", 1)[0]
|
|
409
|
+
if metadata.sklearn_version
|
|
410
|
+
else ""
|
|
411
|
+
)
|
|
412
|
+
if stored_major and stored_major != runtime_major:
|
|
413
|
+
raise ValueError(
|
|
414
|
+
f"model was built with scikit-learn {metadata.sklearn_version}, "
|
|
415
|
+
f"but the runtime has {sklearn.__version__}. "
|
|
416
|
+
"Retrain/export the model for the current major version."
|
|
417
|
+
)
|
|
418
|
+
|
|
419
|
+
try:
|
|
420
|
+
from skops import io as skops_io
|
|
421
|
+
except ImportError as exc:
|
|
422
|
+
raise RuntimeError(
|
|
423
|
+
"skops is required for model loading; install with pip install skops"
|
|
424
|
+
) from exc
|
|
425
|
+
|
|
426
|
+
unknown = list(skops_io.get_untrusted_types(file=path))
|
|
427
|
+
allowed_types = set(DEFAULT_TRUSTED_SKOPS_TYPES).union(trusted_types)
|
|
428
|
+
unresolved = [name for name in unknown if name not in allowed_types]
|
|
429
|
+
if unresolved:
|
|
430
|
+
raise ValueError(
|
|
431
|
+
"model contains unknown serialized types; inspect the .skops file before loading: "
|
|
432
|
+
+ ", ".join(unresolved)
|
|
433
|
+
)
|
|
434
|
+
model = skops_io.load(path, trusted=sorted(allowed_types))
|
|
435
|
+
if not isinstance(model, RandomForestClassifier):
|
|
436
|
+
raise ValueError(
|
|
437
|
+
"unsupported model type in CandidateSuppressor artifact; expected RandomForestClassifier"
|
|
438
|
+
)
|
|
439
|
+
obj = cls(model=model, metadata=metadata)
|
|
440
|
+
obj.feature_names = list(envelope.get("feature_names", [])) or None
|
|
441
|
+
return obj
|