MNT metadata routing: remove `MethodMapping.from_str()` and sort `caller`, `callee` in `MethodPair()` (#28422)

Co-authored-by: Adrin Jalali <adrin.jalali@gmail.com>
Co-authored-by: Guillaume Lemaitre <guillaume@probabl.ai>
This commit is contained in:
Stefanie Senger 2024-05-02 10:13:21 +02:00 committed by GitHub
parent f61dd6cfc0
commit 2bafd7b68f
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
14 changed files with 113 additions and 129 deletions

View File

@ -167,9 +167,9 @@ class MetaClassifier(MetaEstimatorMixin, ClassifierMixin, BaseEstimator):
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(callee="fit", caller="fit")
.add(callee="predict", caller="predict")
.add(callee="score", caller="score"),
.add(caller="fit", callee="fit")
.add(caller="predict", callee="predict")
.add(caller="score", callee="score"),
)
return router
@ -356,9 +356,9 @@ class RouterConsumerClassifier(MetaEstimatorMixin, ClassifierMixin, BaseEstimato
.add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(callee="fit", caller="fit")
.add(callee="predict", caller="predict")
.add(callee="score", caller="score"),
.add(caller="fit", callee="fit")
.add(caller="predict", callee="predict")
.add(caller="score", callee="score"),
)
)
return router
@ -488,16 +488,16 @@ class SimplePipeline(ClassifierMixin, BaseEstimator):
# The metadata is routed such that it retraces how
# `SimplePipeline` internally calls the transformer's `fit` and
# `transform` methods in its own methods (`fit` and `predict`).
.add(callee="fit", caller="fit")
.add(callee="transform", caller="fit")
.add(callee="transform", caller="predict"),
.add(caller="fit", callee="fit")
.add(caller="fit", callee="transform")
.add(caller="predict", callee="transform"),
)
# We add the routing for the classifier.
.add(
classifier=self.classifier,
method_mapping=MethodMapping()
.add(callee="fit", caller="fit")
.add(callee="predict", caller="predict"),
.add(caller="fit", callee="fit")
.add(caller="predict", callee="predict"),
)
)
return router
@ -612,7 +612,7 @@ class MetaRegressor(MetaEstimatorMixin, RegressorMixin, BaseEstimator):
def get_metadata_routing(self):
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.estimator,
method_mapping=MethodMapping().add(callee="fit", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
return router
@ -651,7 +651,7 @@ class WeightedMetaRegressor(MetaEstimatorMixin, RegressorMixin, BaseEstimator):
.add_self_request(self)
.add(
estimator=self.estimator,
method_mapping=MethodMapping().add(callee="fit", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
)
return router
@ -692,7 +692,7 @@ for w in record:
print(w.message)
# %%
# In the end, we disable the configuration flag for metadata routing:
# At the end we disable the configuration flag for metadata routing:
set_config(enable_metadata_routing=False)

View File

@ -523,11 +523,11 @@ class CalibratedClassifierCV(ClassifierMixin, MetaEstimatorMixin, BaseEstimator)
.add_self_request(self)
.add(
estimator=self._get_estimator(),
method_mapping=MethodMapping().add(callee="fit", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
.add(
splitter=self.cv,
method_mapping=MethodMapping().add(callee="split", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="split"),
)
)
return router

View File

@ -513,8 +513,8 @@ class SelectFromModel(MetaEstimatorMixin, SelectorMixin, BaseEstimator):
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(callee="partial_fit", caller="partial_fit")
.add(callee="fit", caller="fit"),
.add(caller="partial_fit", callee="partial_fit")
.add(caller="fit", callee="fit"),
)
return router

View File

@ -1860,7 +1860,7 @@ class LinearModelCV(MultiOutputMixin, LinearModel, ABC):
.add_self_request(self)
.add(
splitter=check_cv(self.cv),
method_mapping=MethodMapping().add(callee="split", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="split"),
)
)
return router

View File

@ -1821,7 +1821,7 @@ class LarsCV(Lars):
"""
router = MetadataRouter(owner=self.__class__.__name__).add(
splitter=check_cv(self.cv),
method_mapping=MethodMapping().add(callee="split", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="split"),
)
return router

View File

@ -2166,13 +2166,13 @@ class LogisticRegressionCV(LogisticRegression, LinearClassifierMixin, BaseEstima
.add_self_request(self)
.add(
splitter=self.cv,
method_mapping=MethodMapping().add(callee="split", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="split"),
)
.add(
scorer=self._get_scorer(),
method_mapping=MethodMapping()
.add(callee="score", caller="score")
.add(callee="score", caller="fit"),
.add(caller="score", callee="score")
.add(caller="fit", callee="score"),
)
)
return router

View File

@ -1116,6 +1116,6 @@ class OrthogonalMatchingPursuitCV(RegressorMixin, LinearModel):
router = MetadataRouter(owner=self.__class__.__name__).add(
splitter=self.cv,
method_mapping=MethodMapping().add(callee="split", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="split"),
)
return router

View File

@ -32,6 +32,7 @@ from ..utils._response import _get_response_values
from ..utils.metadata_routing import (
MetadataRequest,
MetadataRouter,
MethodMapping,
_MetadataRequester,
_raise_for_params,
_routing_enabled,
@ -188,7 +189,8 @@ class _MultimetricScorer:
routing information.
"""
return MetadataRouter(owner=self.__class__.__name__).add(
**self._scorers, method_mapping="score"
**self._scorers,
method_mapping=MethodMapping().add(caller="score", callee="score"),
)

View File

@ -61,7 +61,7 @@ from sklearn.utils._testing import (
assert_array_equal,
ignore_warnings,
)
from sklearn.utils.metadata_routing import MetadataRouter
from sklearn.utils.metadata_routing import MetadataRouter, MethodMapping
REGRESSION_SCORERS = [
"d2_absolute_error_score",
@ -1233,7 +1233,8 @@ def test_scorer_metadata_request(name):
# make sure putting the scorer in a router doesn't request anything by
# default
router = MetadataRouter(owner="test").add(
method_mapping="score", scorer=get_scorer(name)
scorer=get_scorer(name),
method_mapping=MethodMapping().add(caller="score", callee="score"),
)
# make sure `sample_weight` is refused if passed.
with pytest.raises(TypeError, match="got unexpected argument"):
@ -1244,7 +1245,8 @@ def test_scorer_metadata_request(name):
# make sure putting weighted_scorer in a router requests sample_weight
router = MetadataRouter(owner="test").add(
scorer=weighted_scorer, method_mapping="score"
scorer=weighted_scorer,
method_mapping=MethodMapping().add(caller="score", callee="score"),
)
router.validate_metadata(params={"sample_weight": 1}, method="score")
routed_params = router.route_params(params={"sample_weight": 1}, caller="score")

View File

@ -619,8 +619,8 @@ class OneVsRestClassifier(
.add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(callee="fit", caller="fit")
.add(callee="partial_fit", caller="partial_fit"),
.add(caller="fit", callee="fit")
.add(caller="partial_fit", callee="partial_fit"),
)
)
return router
@ -1018,8 +1018,8 @@ class OneVsOneClassifier(MetaEstimatorMixin, ClassifierMixin, BaseEstimator):
.add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(callee="fit", caller="fit")
.add(callee="partial_fit", caller="partial_fit"),
.add(caller="fit", callee="fit")
.add(caller="partial_fit", callee="partial_fit"),
)
)
return router
@ -1264,6 +1264,6 @@ class OutputCodeClassifier(MetaEstimatorMixin, ClassifierMixin, BaseEstimator):
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.estimator,
method_mapping=MethodMapping().add(callee="fit", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
return router

View File

@ -334,8 +334,8 @@ class _MultiOutputEstimator(MetaEstimatorMixin, BaseEstimator, metaclass=ABCMeta
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(callee="partial_fit", caller="partial_fit")
.add(callee="fit", caller="fit"),
.add(caller="partial_fit", callee="partial_fit")
.add(caller="fit", callee="fit"),
)
return router
@ -1096,7 +1096,7 @@ class ClassifierChain(MetaEstimatorMixin, ClassifierMixin, _BaseChain):
"""
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.base_estimator,
method_mapping=MethodMapping().add(callee="fit", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
return router
@ -1245,7 +1245,7 @@ class RegressorChain(MetaEstimatorMixin, RegressorMixin, _BaseChain):
"""
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.base_estimator,
method_mapping=MethodMapping().add(callee="fit", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
return router

View File

@ -19,6 +19,7 @@ from sklearn.utils._metadata_requests import (
)
from sklearn.utils.metadata_routing import (
MetadataRouter,
MethodMapping,
process_routing,
)
from sklearn.utils.multiclass import _check_partial_fit_first_call
@ -418,7 +419,8 @@ class MetaRegressor(MetaEstimatorMixin, RegressorMixin, BaseEstimator):
def get_metadata_routing(self):
router = MetadataRouter(owner=self.__class__.__name__).add(
estimator=self.estimator, method_mapping="one-to-one"
estimator=self.estimator,
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
return router
@ -447,7 +449,12 @@ class WeightedMetaRegressor(MetaEstimatorMixin, RegressorMixin, BaseEstimator):
router = (
MetadataRouter(owner=self.__class__.__name__)
.add_self_request(self)
.add(estimator=self.estimator, method_mapping="one-to-one")
.add(
estimator=self.estimator,
method_mapping=MethodMapping()
.add(caller="fit", callee="fit")
.add(caller="predict", callee="predict"),
)
)
return router
@ -472,7 +479,10 @@ class WeightedMetaClassifier(MetaEstimatorMixin, ClassifierMixin, BaseEstimator)
router = (
MetadataRouter(owner=self.__class__.__name__)
.add_self_request(self)
.add(estimator=self.estimator, method_mapping="fit")
.add(
estimator=self.estimator,
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
)
return router
@ -494,5 +504,8 @@ class MetaTransformer(MetaEstimatorMixin, TransformerMixin, BaseEstimator):
def get_metadata_routing(self):
return MetadataRouter(owner=self.__class__.__name__).add(
transformer=self.transformer, method_mapping="one-to-one"
transformer=self.transformer,
method_mapping=MethodMapping()
.add(caller="fit", callee="fit")
.add(caller="transform", callee="transform"),
)

View File

@ -114,11 +114,16 @@ class SimplePipeline(BaseEstimator):
router.add(
**{f"step_{i}": step},
method_mapping=MethodMapping()
.add(callee="fit", caller="fit")
.add(callee="transform", caller="fit")
.add(callee="transform", caller="predict"),
.add(caller="fit", callee="fit")
.add(caller="fit", callee="transform")
.add(caller="predict", callee="transform"),
)
router.add(predictor=self.steps[-1], method_mapping="one-to-one")
router.add(
predictor=self.steps[-1],
method_mapping=MethodMapping()
.add(caller="fit", callee="fit")
.add(caller="predict", callee="predict"),
)
return router
@ -150,7 +155,10 @@ def test_assert_request_is_empty():
assert_request_is_empty(
MetadataRouter(owner="test")
.add_self_request(WeightedMetaRegressor(estimator=None))
.add(method_mapping="fit", estimator=ConsumingRegressor())
.add(
estimator=ConsumingRegressor(),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
)
@ -677,13 +685,13 @@ def test_estimator_warnings():
MetadataRequest(owner="test"),
"{}",
),
(MethodMapping.from_str("score"), "[{'callee': 'score', 'caller': 'score'}]"),
(
MetadataRouter(owner="test").add(
method_mapping="predict", estimator=ConsumingRegressor()
estimator=ConsumingRegressor(),
method_mapping=MethodMapping().add(caller="predict", callee="predict"),
),
(
"{'estimator': {'mapping': [{'callee': 'predict', 'caller':"
"{'estimator': {'mapping': [{'caller': 'predict', 'callee':"
" 'predict'}], 'router': {'fit': {'sample_weight': None, 'metadata':"
" None}, 'partial_fit': {'sample_weight': None, 'metadata': None},"
" 'predict': {'sample_weight': None, 'metadata': None}, 'score':"
@ -702,24 +710,17 @@ def test_string_representations(obj, string):
(
MethodMapping(),
"add",
{"callee": "invalid", "caller": "fit"},
{"caller": "fit", "callee": "invalid"},
ValueError,
"Given callee",
),
(
MethodMapping(),
"add",
{"callee": "fit", "caller": "invalid"},
{"caller": "invalid", "callee": "fit"},
ValueError,
"Given caller",
),
(
MethodMapping,
"from_str",
{"route": "invalid"},
ValueError,
"route should be 'one-to-one' or a single method!",
),
(
MetadataRouter(owner="test"),
"add_self_request",
@ -749,16 +750,17 @@ def test_methodmapping():
)
mm_list = list(mm)
assert mm_list[0] == ("transform", "fit")
assert mm_list[0] == ("fit", "transform")
assert mm_list[1] == ("fit", "fit")
mm = MethodMapping.from_str("one-to-one")
mm = MethodMapping()
for method in METHODS:
mm.add(caller=method, callee=method)
assert MethodPair(method, method) in mm._routes
assert len(mm._routes) == len(METHODS)
mm = MethodMapping.from_str("score")
assert repr(mm) == "[{'callee': 'score', 'caller': 'score'}]"
mm = MethodMapping().add(caller="score", callee="score")
assert repr(mm) == "[{'caller': 'score', 'callee': 'score'}]"
def test_metadatarouter_add_self_request():
@ -793,12 +795,12 @@ def test_metadatarouter_add_self_request():
def test_metadata_routing_add():
# adding one with a string `method_mapping`
router = MetadataRouter(owner="test").add(
method_mapping="fit",
est=ConsumingRegressor().set_fit_request(sample_weight="weights"),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
assert (
str(router)
== "{'est': {'mapping': [{'callee': 'fit', 'caller': 'fit'}], 'router': {'fit':"
== "{'est': {'mapping': [{'caller': 'fit', 'callee': 'fit'}], 'router': {'fit':"
" {'sample_weight': 'weights', 'metadata': None}, 'partial_fit':"
" {'sample_weight': None, 'metadata': None}, 'predict': {'sample_weight':"
" None, 'metadata': None}, 'score': {'sample_weight': None, 'metadata':"
@ -807,12 +809,12 @@ def test_metadata_routing_add():
# adding one with an instance of MethodMapping
router = MetadataRouter(owner="test").add(
method_mapping=MethodMapping().add(callee="score", caller="fit"),
method_mapping=MethodMapping().add(caller="fit", callee="score"),
est=ConsumingRegressor().set_score_request(sample_weight=True),
)
assert (
str(router)
== "{'est': {'mapping': [{'callee': 'score', 'caller': 'fit'}], 'router':"
== "{'est': {'mapping': [{'caller': 'fit', 'callee': 'score'}], 'router':"
" {'fit': {'sample_weight': None, 'metadata': None}, 'partial_fit':"
" {'sample_weight': None, 'metadata': None}, 'predict': {'sample_weight':"
" None, 'metadata': None}, 'score': {'sample_weight': True, 'metadata':"
@ -829,17 +831,17 @@ def test_metadata_routing_get_param_names():
)
)
.add(
method_mapping="fit",
trs=ConsumingTransformer().set_fit_request(
sample_weight="transform_weights"
),
method_mapping=MethodMapping().add(caller="fit", callee="fit"),
)
)
assert (
str(router)
== "{'$self_request': {'fit': {'sample_weight': 'self_weights'}, 'score':"
" {'sample_weight': None}}, 'trs': {'mapping': [{'callee': 'fit', 'caller':"
" {'sample_weight': None}}, 'trs': {'mapping': [{'caller': 'fit', 'callee':"
" 'fit'}], 'router': {'fit': {'sample_weight': 'transform_weights',"
" 'metadata': None}, 'transform': {'sample_weight': None, 'metadata': None},"
" 'inverse_transform': {'sample_weight': None, 'metadata': None}}}}"

View File

@ -693,19 +693,18 @@ RouterMappingPair = namedtuple("RouterMappingPair", ["mapping", "router"])
# A namedtuple storing a single method route. A collection of these namedtuples
# is stored in a MetadataRouter.
MethodPair = namedtuple("MethodPair", ["callee", "caller"])
MethodPair = namedtuple("MethodPair", ["caller", "callee"])
class MethodMapping:
"""Stores the mapping between callee and caller methods for a router.
"""Stores the mapping between caller and callee methods for a router.
This class is primarily used in a ``get_metadata_routing()`` of a router
object when defining the mapping between a sub-object (a sub-estimator or a
scorer) to the router's methods. It stores a collection of ``Route``
namedtuples.
scorer) to the router's methods. It stores a collection of namedtuples.
Iterating through an instance of this class will yield named
``MethodPair(callee, caller)`` tuples.
``MethodPair(caller, callee)`` tuples.
.. versionadded:: 1.3
"""
@ -716,33 +715,34 @@ class MethodMapping:
def __iter__(self):
return iter(self._routes)
def add(self, *, callee, caller):
def add(self, *, caller, callee):
"""Add a method mapping.
Parameters
----------
callee : str
Child object's method name. This method is called in ``caller``.
caller : str
Parent estimator's method name in which the ``callee`` is called.
callee : str
Child object's method name. This method is called in ``caller``.
Returns
-------
self : MethodMapping
Returns self.
"""
if callee not in METHODS:
raise ValueError(
f"Given callee:{callee} is not a valid method. Valid methods are:"
f" {METHODS}"
)
if caller not in METHODS:
raise ValueError(
f"Given caller:{caller} is not a valid method. Valid methods are:"
f" {METHODS}"
)
self._routes.append(MethodPair(callee=callee, caller=caller))
if callee not in METHODS:
raise ValueError(
f"Given callee:{callee} is not a valid method. Valid methods are:"
f" {METHODS}"
)
self._routes.append(MethodPair(caller=caller, callee=callee))
return self
def _serialize(self):
@ -755,38 +755,9 @@ class MethodMapping:
"""
result = list()
for route in self._routes:
result.append({"callee": route.callee, "caller": route.caller})
result.append({"caller": route.caller, "callee": route.callee})
return result
@classmethod
def from_str(cls, route):
"""Construct an instance from a string.
Parameters
----------
route : str
A string representing the mapping, it can be:
- `"one-to-one"`: a one to one mapping for all methods.
- `"method"`: the name of a single method, such as ``fit``,
``transform``, ``score``, etc.
Returns
-------
obj : MethodMapping
A :class:`~sklearn.utils.metadata_routing.MethodMapping` instance
constructed from the given string.
"""
routing = cls()
if route == "one-to-one":
for method in METHODS:
routing.add(callee=method, caller=method)
elif route in METHODS:
routing.add(callee=route, caller=route)
else:
raise ValueError("route should be 'one-to-one' or a single method!")
return routing
def __repr__(self):
return str(self._serialize())
@ -868,10 +839,8 @@ class MetadataRouter:
Parameters
----------
method_mapping : MethodMapping or str
The mapping between the child and the parent's methods. If str, the
output of :func:`~sklearn.utils.metadata_routing.MethodMapping.from_str`
is used.
method_mapping : MethodMapping
The mapping between the child and the parent's methods.
**objs : dict
A dictionary of objects from which metadata is extracted by calling
@ -882,10 +851,7 @@ class MetadataRouter:
self : MetadataRouter
Returns `self`.
"""
if isinstance(method_mapping, str):
method_mapping = MethodMapping.from_str(method_mapping)
else:
method_mapping = deepcopy(method_mapping)
method_mapping = deepcopy(method_mapping)
for name, obj in objs.items():
self._route_mappings[name] = RouterMappingPair(
@ -916,7 +882,7 @@ class MetadataRouter:
res = res | self._self_request.consumes(method=method, params=params)
for _, route_mapping in self._route_mappings.items():
for callee, caller in route_mapping.mapping:
for caller, callee in route_mapping.mapping:
if caller == method:
res = res | route_mapping.router.consumes(
method=callee, params=params
@ -959,7 +925,7 @@ class MetadataRouter:
)
for name, route_mapping in self._route_mappings.items():
for callee, caller in route_mapping.mapping:
for caller, callee in route_mapping.mapping:
if caller == method:
res = res.union(
route_mapping.router._get_param_names(
@ -1065,7 +1031,7 @@ class MetadataRouter:
router, mapping = route_mapping.router, route_mapping.mapping
res[name] = Bunch()
for _callee, _caller in mapping:
for _caller, _callee in mapping:
if _caller == caller:
res[name][_callee] = router._route_params(
params=params,
@ -1127,12 +1093,11 @@ class MetadataRouter:
def __iter__(self):
if self._self_request:
yield (
"$self_request",
RouterMappingPair(
mapping=MethodMapping.from_str("one-to-one"),
router=self._self_request,
),
method_mapping = MethodMapping()
for method in METHODS:
method_mapping.add(caller=method, callee=method)
yield "$self_request", RouterMappingPair(
mapping=method_mapping, router=self._self_request
)
for name, route_mapping in self._route_mappings.items():
yield (name, route_mapping)