Check sklearn's clone() preserves init parameters.
Delegates the shallow-param and nested-estimator contract to the shared,
family-agnostic check_clone_preserves_params (which uses _safe_equal
and so tolerates DataFrame/array-valued params whose == is not a bool).
Adds the forecaster-specific check for (name, estimator, columns)
3-tuples (e.g. ColumnForecaster): the trailing columns element must
survive clone unchanged.
Parameters
| Name |
Type |
Description |
Default |
forecaster
|
BaseForecaster
|
|
required
|
Raises
| Type |
Description |
AssertionError
|
If cloned forecaster has different parameters
|
Source Code
View on GitHub
Source code in src/yohou/testing/forecaster.py
| def check_clone_preserves_forecaster_params(forecaster) -> None:
"""Check sklearn's clone() preserves init parameters.
Delegates the shallow-param and nested-estimator contract to the shared,
family-agnostic ``check_clone_preserves_params`` (which uses ``_safe_equal``
and so tolerates DataFrame/array-valued params whose ``==`` is not a bool).
Adds the forecaster-specific check for ``(name, estimator, columns)``
3-tuples (e.g. ``ColumnForecaster``): the trailing ``columns`` element must
survive clone unchanged.
Parameters
----------
forecaster : BaseForecaster
Forecaster instance
Raises
------
AssertionError
If cloned forecaster has different parameters
"""
check_clone_preserves_params(forecaster)
forecaster_clone = clone(forecaster)
original_params = forecaster.get_params(deep=False)
cloned_params = forecaster_clone.get_params(deep=False)
for key, orig_val in original_params.items():
cloned_val = cloned_params[key]
if not (isinstance(orig_val, list) and orig_val and isinstance(orig_val[0], tuple)):
continue
for i, (orig_item, cloned_item) in enumerate(zip(orig_val, cloned_val, strict=True)):
if not (isinstance(orig_item, tuple) and isinstance(cloned_item, tuple)):
continue
if len(orig_item) == 3 and len(cloned_item) == 3:
orig_cols, cloned_cols = orig_item[-1], cloned_item[-1]
assert _safe_equal(orig_cols, cloned_cols), (
f"Parameter {key}[{i}] columns: {cloned_cols!r} != {orig_cols!r}"
)
assert forecaster_clone is not forecaster, "clone() should create new instance"
|