Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion src/pyrecest/utils/association_features.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,7 +403,10 @@ def _finite_feature_plane(values: Any, feature_name: str) -> Any:


def _flatten_prediction_features(features: Any) -> tuple[Any, tuple[int, ...]]:
features = asarray(features, dtype=float64)
features = _as_real_numeric_backend_array(
features,
message="features must contain real numeric values",
)
if features.ndim == 0:
raise ValueError("features must be at least one-dimensional")
if features.ndim == 1:
Expand Down
31 changes: 31 additions & 0 deletions tests/utils/test_calibrated_association_feature_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""Regression tests for calibrated association feature validation."""

import pytest

from pyrecest.backend import array
from pyrecest.utils import CalibratedPairwiseAssociationModel


class _RecordingPredictProbaModel:
classes_ = array([0, 1])

def __init__(self):
self.called = False

def predict_proba(self, features):
self.called = True
return array([[0.25, 0.75]])


def test_predict_proba_rejects_complex_direct_features_before_model_call():
model = _RecordingPredictProbaModel()
calibrated_model = CalibratedPairwiseAssociationModel(
model,
feature_names=("distance", "similarity"),
)
complex_features = array([[1.0 + 2.0j, 0.5]])

with pytest.raises(ValueError, match="real numeric"):
calibrated_model.predict_match_probability(complex_features)

assert not model.called
Loading