diff --git a/pretab/preprocessor.py b/pretab/preprocessor.py index fcf228a..a1ab085 100644 --- a/pretab/preprocessor.py +++ b/pretab/preprocessor.py @@ -9,6 +9,7 @@ from .core.exceptions import ( IncompatibleParamsError, + PretabDataError, invalid_param_error, ) from .core.logging import configure_logging, get_logger @@ -360,11 +361,8 @@ def fit(self, X, y=None, embeddings=None): self.embedding_dimensions_ = {} if embeddings is not None: self.embeddings_ = True - if isinstance(embeddings, np.ndarray): - self.embedding_dimensions_["embedding_1"] = embeddings.shape[1] - elif isinstance(embeddings, list): - for i, e in enumerate(embeddings): - self.embedding_dimensions_[f"embedding_{i + 1}"] = e.shape[1] + for i, e in enumerate(self._as_embedding_list(embeddings)): + self.embedding_dimensions_[f"embedding_{i + 1}"] = e.shape[1] numerical_features, categorical_features = self._detect_column_types(X) transformers = [] @@ -476,14 +474,53 @@ def transform(self, X, embeddings=None, return_array=False): "Fix: configure an embedding feature in feature_preprocessing before " "passing embeddings to transform, or omit the embeddings argument." ) - if isinstance(embeddings, np.ndarray): - transformed_dict["embedding_1"] = embeddings.astype(np.float32) - elif isinstance(embeddings, list): - for idx, e in enumerate(embeddings): - transformed_dict[f"embedding_{idx + 1}"] = e.astype(np.float32) + for idx, e in enumerate(self._validated_embeddings(embeddings, len(X))): + transformed_dict[f"embedding_{idx + 1}"] = e.astype(np.float32) return transformed_dict + @staticmethod + def _as_embedding_list(embeddings) -> list: + """Normalize a single array or a list of arrays to a list of arrays.""" + if isinstance(embeddings, np.ndarray): + return [embeddings] + return list(embeddings) + + def _validated_embeddings(self, embeddings, n_samples: int) -> list: + """Check supplied embeddings against what ``fit`` recorded. + + ``fit`` stores ``embedding_dimensions_`` but nothing used to read it back, + so a mismatched array was accepted silently -- including one with the + wrong number of rows, which produced a result dict whose blocks had + different heights. + """ + arrays = self._as_embedding_list(embeddings) + expected = list(self.embedding_dimensions_.items()) + + if len(arrays) != len(expected): + raise PretabDataError( + f"Expected {len(expected)} embedding array(s) as seen during fit, " + f"got {len(arrays)}." + ) + + for idx, (array, (name, dim)) in enumerate(zip(arrays, expected, strict=True)): + array = np.asarray(array) + if array.ndim != 2: + raise PretabDataError( + f"{name} must be a 2D array, got {array.ndim} dimension(s)." + ) + if array.shape[1] != dim: + raise PretabDataError( + f"{name} has {array.shape[1]} columns, but {dim} were seen during fit." + ) + if array.shape[0] != n_samples: + raise PretabDataError( + f"{name} has {array.shape[0]} rows, but X has {n_samples}." + ) + arrays[idx] = array + + return arrays + def fit_transform(self, X, y=None, embeddings=None, return_array=False): """ Convenience method that fits the preprocessor and transforms the data. diff --git a/tests/test_embeddings.py b/tests/test_embeddings.py new file mode 100644 index 0000000..5aeded2 --- /dev/null +++ b/tests/test_embeddings.py @@ -0,0 +1,94 @@ +"""External embedding arrays must be validated against what ``fit`` recorded. + +``Preprocessor.fit`` populates ``embedding_dimensions_`` but nothing used to read +it back, so ``transform`` accepted any array -- including one with the wrong row +count, which produced a result dict whose blocks had different heights. +""" + +from typing import cast + +import numpy as np +import pandas as pd +import pytest + +from pretab.core.exceptions import IncompatibleParamsError, PretabDataError +from pretab.preprocessor import Preprocessor + + +@pytest.fixture +def data(): + rng = np.random.default_rng(0) + frame = pd.DataFrame({"a": rng.normal(size=100), "c": rng.choice(list("xyz"), 100)}) + return frame, rng.normal(size=100), rng + + +def _fitted(frame, y, embeddings): + return Preprocessor(numerical_method="minmax").fit(frame, y, embeddings=embeddings) + + +def test_matching_embeddings_pass_through(data): + frame, y, rng = data + pre = _fitted(frame, y, rng.random((100, 8))) + + out = cast("dict[str, np.ndarray]", pre.transform(frame, embeddings=rng.random((100, 8)))) + + assert out["embedding_1"].shape == (100, 8) + assert out["embedding_1"].dtype == np.float32 + + +def test_fit_records_the_dimensions(data): + frame, y, rng = data + assert _fitted(frame, y, rng.random((100, 8))).embedding_dimensions_ == {"embedding_1": 8} + + +def test_wrong_width_is_rejected(data): + frame, y, rng = data + pre = _fitted(frame, y, rng.random((100, 8))) + + with pytest.raises(PretabDataError, match="has 3 columns, but 8 were seen during fit"): + pre.transform(frame, embeddings=rng.random((100, 3))) + + +def test_wrong_row_count_is_rejected(data): + frame, y, rng = data + pre = _fitted(frame, y, rng.random((100, 8))) + + with pytest.raises(PretabDataError, match="has 7 rows, but X has 100"): + pre.transform(frame, embeddings=rng.random((7, 8))) + + +def test_one_dimensional_embedding_is_rejected(data): + frame, y, rng = data + pre = _fitted(frame, y, rng.random((100, 8))) + + with pytest.raises(PretabDataError, match="must be a 2D array"): + pre.transform(frame, embeddings=rng.random(100)) + + +def test_wrong_number_of_arrays_is_rejected(data): + frame, y, rng = data + pre = _fitted(frame, y, [rng.random((100, 4)), rng.random((100, 5))]) + + with pytest.raises(PretabDataError, match="Expected 2 embedding array"): + pre.transform(frame, embeddings=[rng.random((100, 4))]) + + +def test_embedding_list_round_trips(data): + frame, y, rng = data + pre = _fitted(frame, y, [rng.random((100, 4)), rng.random((100, 5))]) + + out = cast( + "dict[str, np.ndarray]", + pre.transform(frame, embeddings=[rng.random((100, 4)), rng.random((100, 5))]), + ) + + assert out["embedding_1"].shape == (100, 4) + assert out["embedding_2"].shape == (100, 5) + + +def test_unexpected_embeddings_still_rejected(data): + frame, y, rng = data + pre = Preprocessor(numerical_method="minmax").fit(frame, y) + + with pytest.raises(IncompatibleParamsError): + pre.transform(frame, embeddings=rng.random((100, 8)))