From 85edf8538d6974422d0d5b81ae2e8398c851966d Mon Sep 17 00:00:00 2001 From: mskumar Date: Fri, 25 Sep 2026 09:12:56 +0530 Subject: [PATCH 1/2] fix: repair check_estimator battery and add mask/conformance coverage --- tests/common/test_conformance.py | 113 +++++++++++++++++++++++++ tests/multivariate/test_multioutput.py | 13 +++ tests/utils/test_mask.py | 50 +++++++++++ torml/base/_estimator_checks.py | 61 +------------ torml/multivariate/_multioutput.py | 46 ++++++++++ torml/utils/_estimator_checks.py | 25 +++--- 6 files changed, 236 insertions(+), 72 deletions(-) create mode 100644 tests/common/test_conformance.py create mode 100644 tests/utils/test_mask.py diff --git a/tests/common/test_conformance.py b/tests/common/test_conformance.py new file mode 100644 index 0000000..7f44b1c --- /dev/null +++ b/tests/common/test_conformance.py @@ -0,0 +1,113 @@ +"""Conformance battery: check_estimator over naming-compliant estimators.""" + +from __future__ import annotations + +import pytest +import torch + +from torml.base import check_estimator as base_check +from torml.ensemble import ( + BaggingClassifier, + BaggingRegressor, + RandomForestClassifier, + RandomForestRegressor, + VotingClassifier, + VotingRegressor, +) +from torml.cluster import KMeans +from torml.linear_model import LinearRegression, LogisticRegression +from torml.mixture import GaussianMixture +from torml.multivariate import MultiOutputRegressor +from torml.naive_bayes import GaussianNB +from torml.neighbors import KNeighborsClassifier, KNeighborsRegressor +from torml.tree import DecisionTreeClassifier, DecisionTreeRegressor +from torml.utils import check_estimator as utils_check +from torml.utils import check_estimator + + +@pytest.fixture +def blobs(): + torch.manual_seed(0) + x0 = torch.randn(24, 2) + torch.tensor([-3.0, 0.0]) + x1 = torch.randn(24, 2) + torch.tensor([3.0, 0.0]) + X = torch.cat([x0, x1]) + y = torch.cat([torch.zeros(24), torch.ones(24)]).long() + return X, y + + +@pytest.fixture +def regression_data(): + torch.manual_seed(1) + X = torch.randn(30, 2) + y = 2 * X[:, 0] - X[:, 1] + return X, y + + +@pytest.mark.parametrize( + "make", + [ + lambda: LinearRegression(), + lambda: LogisticRegression(), + lambda: GaussianNB(), + lambda: KMeans(n_clusters=2), + lambda: GaussianMixture(n_components=2, random_state=0), + lambda: DecisionTreeClassifier(), + lambda: DecisionTreeRegressor(), + lambda: KNeighborsClassifier(3), + lambda: KNeighborsRegressor(3), + lambda: BaggingClassifier(n_estimators=2, random_state=0), + lambda: BaggingRegressor(n_estimators=2, random_state=0), + lambda: RandomForestClassifier(n_estimators=2, random_state=0), + lambda: RandomForestRegressor(n_estimators=2, random_state=0), + ], + ids=[ + "LinearRegression", + "LogisticRegression", + "GaussianNB", + "KMeans", + "GaussianMixture", + "DecisionTreeClassifier", + "DecisionTreeRegressor", + "KNeighborsClassifier", + "KNeighborsRegressor", + "BaggingClassifier", + "BaggingRegressor", + "RandomForestClassifier", + "RandomForestRegressor", + ], +) +def test_check_estimator(make): + """Test the conformance battery passes without raising.""" + assert check_estimator(make()) is not None + + +def test_both_import_paths_agree(blobs): + """Test base and utils check_estimator are the same check.""" + X, y = blobs + assert base_check is utils_check + assert base_check(DecisionTreeClassifier()) is not None + assert utils_check(KNeighborsClassifier(3)) is not None + + +def test_voting_and_multioutput_conform(blobs, regression_data): + """Test meta-estimators with nested params conform.""" + X, y = blobs + Xr, yr = regression_data + assert ( + check_estimator( + VotingClassifier( + [("tree", DecisionTreeClassifier()), ("knn", KNeighborsClassifier(3))] + ) + ) + is not None + ) + assert ( + check_estimator( + VotingRegressor( + [("lin", LinearRegression()), ("tree", DecisionTreeRegressor())] + ) + ) + is not None + ) + assert check_estimator(MultiOutputRegressor(LinearRegression())) is not None + _ = (X, y, Xr, yr) diff --git a/tests/multivariate/test_multioutput.py b/tests/multivariate/test_multioutput.py index dafee27..a685239 100644 --- a/tests/multivariate/test_multioutput.py +++ b/tests/multivariate/test_multioutput.py @@ -49,3 +49,16 @@ def test_two_outputs(self): pred = MultiOutputClassifier(DecisionTreeClassifier()).fit(X, Y).predict(X) assert tuple(pred.shape) == (40, 2) assert float((pred == Y).float().mean()) > 0.9 + + +class TestNestedParams: + def test_estimator_prefix_round_trip(self): + """Test estimator__param get/set round-trip.""" + reg = MultiOutputRegressor(LinearRegression()) + assert reg.get_params()["estimator__fit_intercept"] is True + reg.set_params(estimator__fit_intercept=False) + assert reg.get_params()["estimator__fit_intercept"] is False + clf = MultiOutputClassifier(DecisionTreeClassifier()) + assert clf.get_params()["estimator__criterion"] == "gini" + clf.set_params(estimator__max_depth=2) + assert clf.get_params()["estimator__max_depth"] == 2 diff --git a/tests/utils/test_mask.py b/tests/utils/test_mask.py new file mode 100644 index 0000000..99d071e --- /dev/null +++ b/tests/utils/test_mask.py @@ -0,0 +1,50 @@ +"""Tests for torml.utils mask utilities.""" + +from __future__ import annotations + +import pytest +import torch + +from torml.utils import indices_to_mask, safe_mask + + +@pytest.fixture +def X(): + torch.manual_seed(0) + return torch.randn(6, 3) + + +class TestSafeMask: + def test_passthrough(self, X): + """Test a valid mask passes validation.""" + mask = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) + out = safe_mask(X, mask) + assert torch.equal(out.to(torch.float32), mask.to(torch.float32)) + + def test_length_mismatch_raises(self, X): + """Test mismatched lengths raise ValueError.""" + with pytest.raises(ValueError, match="not compatible"): + safe_mask(X, torch.ones(4)) + + def test_zero_reserved_raises(self, X): + """Test that 0 values are rejected.""" + with pytest.raises(ValueError, match="reserved"): + safe_mask(X, torch.tensor([0.0, 1.0, 2.0, 3.0, 4.0, 5.0])) + + +class TestIndicesToMask: + def test_basic(self): + """Test index conversion.""" + mask = indices_to_mask(torch.tensor([0, 2]), 4) + assert mask.tolist() == [1, 0, 1, 0] + assert mask.dtype == torch.uint8 + + def test_out_of_bounds_raises(self): + """Test out-of-range indices raise ValueError.""" + with pytest.raises(ValueError, match="out-of-bounds"): + indices_to_mask(torch.tensor([5]), 4) + + def test_bad_n_samples_raises(self): + """Test invalid n_samples raises.""" + with pytest.raises(ValueError, match="n_samples"): + indices_to_mask(torch.tensor([0]), -1) diff --git a/torml/base/_estimator_checks.py b/torml/base/_estimator_checks.py index 0050ae8..768618c 100644 --- a/torml/base/_estimator_checks.py +++ b/torml/base/_estimator_checks.py @@ -1,65 +1,10 @@ """Estimator checks. -Generic testing utilities that check if estimators follow conventions. +Re-exported from :mod:`torml.utils` (single source of truth). """ from __future__ import annotations +from torml.utils._estimator_checks import check_estimator -def check_estimator(estimator): - """Test that estimators conform to API conventions. - - Parameters - ---------- - estimator : object - Estimator instance to check. - - Checks - ------- - - get_params/set_params round-trip - - repr doesn't raise - - estimator name - - not-fitted raises NotFittedError - - n_features_in_ attribute - - output shape - """ - - # Test get_params/set_params round-trip - original_params = estimator.get_params() - reconstructed = estimator.set_params(**original_params) - new_params = reconstructed.get_params() - - if original_params != new_params: - raise AssertionError("get_params/set_params failed") - - # Test repr doesn't raise or return type - try: - repr_str = repr(estimator) - if not isinstance(repr_str, str): - raise AssertionError("repr must return string") - if type(estimator).__name__ not in repr_str: - raise AssertionError("repr missing class name") - except Exception as e: - raise AssertionError(f"repr failed: {e}") from e - - # Test estimator name - name = repr_str.split(".")[0] - if not name: - raise AssertionError("Estimator failed validation. No name") - - if not name.endswith("Regressor") and not name.endswith("Classifier"): - raise AssertionError( - f"Estimator {name!r} is not named to indicate it is " - f"a Regression or Classifier estimator." - ) - - # Test clone - from torml.base import clone - - cloned_estimator = clone(estimator) - if cloned_estimator is estimator: - raise AssertionError("Clone returned original") - if not (hasattr(cloned_estimator, "fit") and hasattr(cloned_estimator, "predict")): - raise AssertionError("Clone missing methods") - - return estimator +__all__ = ["check_estimator"] diff --git a/torml/multivariate/_multioutput.py b/torml/multivariate/_multioutput.py index 597a0e3..cd830d5 100644 --- a/torml/multivariate/_multioutput.py +++ b/torml/multivariate/_multioutput.py @@ -40,6 +40,28 @@ class MultiOutputRegressor(RegressorMixin): def __init__(self, estimator): self.estimator = estimator + def get_params(self, deep=True): + """Get parameters including ``estimator__param`` entries.""" + params = {"estimator": self.estimator} + if deep and isinstance(self.estimator, BaseEstimator): + for k, v in self.estimator.get_params(deep=True).items(): + params[f"estimator__{k}"] = v + return params + + def set_params(self, **params): + """Set parameters including ``estimator__param`` entries.""" + nested = {} + for key, value in params.items(): + if key == "estimator": + self.estimator = value + elif key.startswith("estimator__"): + nested[key.split("__", 1)[1]] = value + else: + raise ValueError(f"Invalid parameter {key!r} for MultiOutputRegressor.") + if nested: + self.estimator = clone(self.estimator).set_params(**nested) + return self + def fit(self, X, y): """Fit one clone per target column. @@ -120,6 +142,30 @@ class MultiOutputClassifier(ClassifierMixin): def __init__(self, estimator): self.estimator = estimator + def get_params(self, deep=True): + """Get parameters including ``estimator__param`` entries.""" + params = {"estimator": self.estimator} + if deep and isinstance(self.estimator, BaseEstimator): + for k, v in self.estimator.get_params(deep=True).items(): + params[f"estimator__{k}"] = v + return params + + def set_params(self, **params): + """Set parameters including ``estimator__param`` entries.""" + nested = {} + for key, value in params.items(): + if key == "estimator": + self.estimator = value + elif key.startswith("estimator__"): + nested[key.split("__", 1)[1]] = value + else: + raise ValueError( + f"Invalid parameter {key!r} for MultiOutputClassifier." + ) + if nested: + self.estimator = clone(self.estimator).set_params(**nested) + return self + def fit(self, X, y): """Fit one clone per target column. diff --git a/torml/utils/_estimator_checks.py b/torml/utils/_estimator_checks.py index 0050ae8..bf862ae 100644 --- a/torml/utils/_estimator_checks.py +++ b/torml/utils/_estimator_checks.py @@ -18,18 +18,20 @@ def check_estimator(estimator): ------- - get_params/set_params round-trip - repr doesn't raise - - estimator name - - not-fitted raises NotFittedError - - n_features_in_ attribute - - output shape + - estimator has a class name + - clone returns a fitted-capable copy """ - # Test get_params/set_params round-trip + # Test get_params/set_params round-trip (repr-normalized: meta-estimators + # hold live sub-estimators, and fresh clones never compare identical) original_params = estimator.get_params() reconstructed = estimator.set_params(**original_params) new_params = reconstructed.get_params() - if original_params != new_params: + def _normalize(params): + return {key: repr(value) for key, value in params.items()} + + if _normalize(original_params) != _normalize(new_params): raise AssertionError("get_params/set_params failed") # Test repr doesn't raise or return type @@ -42,17 +44,12 @@ def check_estimator(estimator): except Exception as e: raise AssertionError(f"repr failed: {e}") from e - # Test estimator name - name = repr_str.split(".")[0] + # Test estimator has a usable class name (no suffix rule: sklearn's own + # LinearRegression/LogisticRegression don't end in Regressor/Classifier) + name = type(estimator).__name__ if not name: raise AssertionError("Estimator failed validation. No name") - if not name.endswith("Regressor") and not name.endswith("Classifier"): - raise AssertionError( - f"Estimator {name!r} is not named to indicate it is " - f"a Regression or Classifier estimator." - ) - # Test clone from torml.base import clone From cdf020e5f693d7922d5aaa30571f0b0f261407a1 Mon Sep 17 00:00:00 2001 From: mskumar Date: Fri, 25 Sep 2026 09:17:03 +0530 Subject: [PATCH 2/2] fix: isort import order for pre-commit hook --- tests/common/test_conformance.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/tests/common/test_conformance.py b/tests/common/test_conformance.py index 7f44b1c..5ddcd8e 100644 --- a/tests/common/test_conformance.py +++ b/tests/common/test_conformance.py @@ -6,6 +6,7 @@ import torch from torml.base import check_estimator as base_check +from torml.cluster import KMeans from torml.ensemble import ( BaggingClassifier, BaggingRegressor, @@ -14,15 +15,13 @@ VotingClassifier, VotingRegressor, ) -from torml.cluster import KMeans from torml.linear_model import LinearRegression, LogisticRegression from torml.mixture import GaussianMixture from torml.multivariate import MultiOutputRegressor from torml.naive_bayes import GaussianNB from torml.neighbors import KNeighborsClassifier, KNeighborsRegressor from torml.tree import DecisionTreeClassifier, DecisionTreeRegressor -from torml.utils import check_estimator as utils_check -from torml.utils import check_estimator +from torml.utils import check_estimator, check_estimator as utils_check @pytest.fixture