Skip to content

Commit 6722001

Browse files
fix(neural): bind quality gates to artifacts
1 parent 3b795d4 commit 6722001

2 files changed

Lines changed: 44 additions & 0 deletions

File tree

daemon/pilot/neural/decoder.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,10 @@ class SSVEPCalibrationArtifact(BaseModel):
8181
coefficients: tuple[tuple[float, ...], ...]
8282
intercepts: tuple[float, ...]
8383
metrics: CalibrationMetrics
84+
minimum_balanced_accuracy: float = Field(default=0.65, ge=0, le=1)
85+
minimum_per_class_recall: float = Field(default=0.50, ge=0, le=1)
86+
maximum_expected_calibration_error: float = Field(default=0.25, ge=0, le=1)
87+
minimum_chance_advantage: float = Field(default=0.10, ge=0, le=1)
8488

8589
@model_validator(mode="after")
8690
def validate_shapes(self) -> Self:
@@ -93,6 +97,8 @@ def validate_shapes(self) -> Self:
9397
raise ValueError("SSVEP frequencies must be unique")
9498
if set(self.class_order) != {target.target_id for target in self.targets}:
9599
raise ValueError("class_order must cover every target exactly once")
100+
if set(self.metrics.per_class_recall) != set(self.class_order):
101+
raise ValueError("per-class recall must cover every calibrated target exactly once")
96102
if len(self.feature_mean) != feature_count or len(self.feature_scale) != feature_count:
97103
raise ValueError("feature normalization shape is invalid")
98104
expected_rows = 1 if class_count == 2 else class_count
@@ -102,6 +108,16 @@ def validate_shapes(self) -> Self:
102108
raise ValueError("classifier feature shape is invalid")
103109
if any(value <= 0 for value in self.feature_scale):
104110
raise ValueError("feature scales must be positive")
111+
required_accuracy = max(
112+
self.minimum_balanced_accuracy,
113+
(1.0 / class_count) + self.minimum_chance_advantage,
114+
)
115+
if self.metrics.balanced_accuracy < required_accuracy:
116+
raise ValueError("artifact balanced accuracy is below its registered threshold")
117+
if min(self.metrics.per_class_recall.values()) < self.minimum_per_class_recall:
118+
raise ValueError("artifact per-class recall is below its registered threshold")
119+
if self.metrics.expected_calibration_error > self.maximum_expected_calibration_error:
120+
raise ValueError("artifact calibration error exceeds its registered threshold")
105121
return self
106122

107123
def content_payload(self) -> bytes:
@@ -338,6 +354,10 @@ def fit(
338354
coefficients=coefficients,
339355
intercepts=tuple(float(value) for value in model.intercept_),
340356
metrics=metrics,
357+
minimum_balanced_accuracy=self._minimum_balanced_accuracy,
358+
minimum_per_class_recall=self._minimum_per_class_recall,
359+
maximum_expected_calibration_error=self._maximum_expected_calibration_error,
360+
minimum_chance_advantage=self._minimum_chance_advantage,
341361
)
342362
calibration_id = hashlib.sha256(base.content_payload()).hexdigest()
343363
artifact = base.model_copy(update={"calibration_id": calibration_id})

daemon/tests/test_neural_decoder.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,30 @@ def test_calibration_requires_registered_advantage_over_chance() -> None:
154154
calibrator._validate_metrics(metrics)
155155

156156

157+
def test_calibration_artifact_binds_acceptance_criteria_and_target_metrics(tmp_path: Path) -> None:
158+
artifact = _artifact()
159+
payload = artifact.model_dump(mode="json")
160+
payload["metrics"]["per_class_recall"] = {
161+
"unrelated-a": 1.0,
162+
"unrelated-b": 1.0,
163+
"unrelated-c": 1.0,
164+
"unrelated-d": 1.0,
165+
}
166+
167+
with pytest.raises(ValueError, match="cover every calibrated target"):
168+
SSVEPCalibrationArtifact.model_validate(payload)
169+
170+
payload = artifact.model_dump(mode="json")
171+
payload["metrics"]["balanced_accuracy"] = 0.2
172+
with pytest.raises(ValueError, match="below its registered threshold"):
173+
SSVEPCalibrationArtifact.model_validate(payload)
174+
175+
artifact.save(tmp_path / "calibration.json")
176+
loaded = SSVEPCalibrationArtifact.load(tmp_path / "calibration.json")
177+
assert loaded.minimum_per_class_recall == 0.5
178+
assert loaded.maximum_expected_calibration_error == 0.25
179+
180+
157181
def test_sidecar_emits_only_signed_derived_intent_and_tracks_dwell() -> None:
158182
artifact = _artifact()
159183
source = SyntheticNeuralSource(target_hz=15, noise_uv=1.0, seed=99)

0 commit comments

Comments
 (0)