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
112 changes: 112 additions & 0 deletions tests/common/test_conformance.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
"""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.cluster import KMeans
from torml.ensemble import (
BaggingClassifier,
BaggingRegressor,
RandomForestClassifier,
RandomForestRegressor,
VotingClassifier,
VotingRegressor,
)
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, check_estimator as utils_check


@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)
13 changes: 13 additions & 0 deletions tests/multivariate/test_multioutput.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
50 changes: 50 additions & 0 deletions tests/utils/test_mask.py
Original file line number Diff line number Diff line change
@@ -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)
61 changes: 3 additions & 58 deletions torml/base/_estimator_checks.py
Original file line number Diff line number Diff line change
@@ -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"]
46 changes: 46 additions & 0 deletions torml/multivariate/_multioutput.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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.

Expand Down
25 changes: 11 additions & 14 deletions torml/utils/_estimator_checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down
Loading