Skip to content

check_clone_preserves_forecaster_params

yohou.testing.check_clone_preserves_forecaster_params(forecaster)

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

Forecaster instance

required

Raises

Type Description
AssertionError

If cloned forecaster has different parameters

Source Code

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"