@@ -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 })
0 commit comments