Skip to content

validate_transformer_data

yohou.utils.validate_transformer_data(transformer, X=None, *, reset=True, inverse=False, X_t=None, X_p=None, observation_horizon=None, stateful=False, **check_params)

validate_transformer_data(
    transformer: BaseActualTransformer,
    X: pl.DataFrame | None = None,
    *,
    reset: Literal[True],
    inverse: bool = False,
    X_t: pl.DataFrame | None = None,
    X_p: pl.DataFrame | None = None,
    observation_horizon: int | None = None,
    stateful: bool = False,
    **check_params,
) -> pl.DataFrame
validate_transformer_data(
    transformer: BaseActualTransformer,
    X: pl.DataFrame | None = None,
    *,
    reset: Literal[False],
    inverse: Literal[True],
    X_t: pl.DataFrame | None = None,
    X_p: pl.DataFrame | None = None,
    observation_horizon: int | None = None,
    stateful: Literal[True],
    **check_params,
) -> tuple[pl.DataFrame, pl.DataFrame]
validate_transformer_data(
    transformer: BaseActualTransformer,
    X: pl.DataFrame | None = None,
    *,
    reset: Literal[False],
    inverse: Literal[True],
    X_t: pl.DataFrame | None = None,
    X_p: pl.DataFrame | None = None,
    observation_horizon: int | None = None,
    stateful: Literal[False] = ...,
    **check_params,
) -> tuple[pl.DataFrame, None]
validate_transformer_data(
    transformer: BaseActualTransformer,
    X: pl.DataFrame | None = None,
    *,
    reset: Literal[False],
    inverse: Literal[False] = ...,
    X_t: pl.DataFrame | None = None,
    X_p: pl.DataFrame | None = None,
    observation_horizon: int | None = None,
    stateful: bool = False,
    **check_params,
) -> pl.DataFrame

Validate data for transformers.

Parameters

Name Type Description Default
transformer BaseActualTransformer

The transformer instance.

required
X DataFrame or None

Input data.

None
reset bool

Whether this is a fit context.

True
inverse bool

Whether this is an inverse transform context.

False
X_t DataFrame or None

Transformed data for inverse transform.

None
X_p DataFrame or None

Previous untransformed data for stateful inverse transform.

None
observation_horizon int or None

Required observation horizon for inverse transform.

None
stateful bool

If True (and inverse=True), X_p is required and guaranteed non-None in return. Use Literal[True] at call site for type narrowing.

False
**check_params dict

Additional validation parameters, consumed only in the transform context (reset=False, inverse=False). Supported keys:

  • check_intervals (bool, default True): run the interval-consistency check on X when X has >= 2 rows.
  • check_continuity (bool, default True): verify X is contiguous with the transformer's observed buffer.
{}

Returns

Type Description
DataFrame or tuple

Validated data. The shape depends on the context (see the overloads):

  • pl.DataFrame: fit context (reset=True) or forward-transform context.
  • tuple[pl.DataFrame, None]: inverse non-stateful context (reset=False, inverse=True, stateful=False).
  • tuple[pl.DataFrame, pl.DataFrame]: inverse stateful context (reset=False, inverse=True, stateful=True).

The first element is always the (transformed) input X; the second element, when present, is X_p.

See Also

Source Code

Source code in src/yohou/utils/validate_data.py
def validate_transformer_data(
    transformer: BaseActualTransformer,
    X: pl.DataFrame | None = None,
    *,
    reset: bool = True,
    inverse: bool = False,
    X_t: pl.DataFrame | None = None,
    X_p: pl.DataFrame | None = None,
    observation_horizon: int | None = None,
    stateful: bool = False,
    **check_params,
) -> pl.DataFrame | tuple[pl.DataFrame, pl.DataFrame | None] | tuple[pl.DataFrame, pl.DataFrame]:
    """Validate data for transformers.

    Parameters
    ----------
    transformer : BaseActualTransformer
        The transformer instance.
    X : pl.DataFrame or None, default=None
        Input data.
    reset : bool, default=True
        Whether this is a fit context.
    inverse : bool, default=False
        Whether this is an inverse transform context.
    X_t : pl.DataFrame or None, default=None
        Transformed data for inverse transform.
    X_p : pl.DataFrame or None, default=None
        Previous untransformed data for stateful inverse transform.
    observation_horizon : int or None, default=None
        Required observation horizon for inverse transform.
    stateful : bool, default=False
        If True (and inverse=True), X_p is required and guaranteed non-None in return.
        Use Literal[True] at call site for type narrowing.
    **check_params : dict
        Additional validation parameters, consumed only in the transform
        context (reset=False, inverse=False). Supported keys:

        - ``check_intervals`` (bool, default True): run the interval-consistency
          check on X when X has >= 2 rows.
        - ``check_continuity`` (bool, default True): verify X is contiguous with
          the transformer's observed buffer.

    Returns
    -------
    pl.DataFrame or tuple
        Validated data. The shape depends on the context (see the overloads):

        - ``pl.DataFrame``: fit context (reset=True) or forward-transform
          context.
        - ``tuple[pl.DataFrame, None]``: inverse non-stateful context
          (reset=False, inverse=True, stateful=False).
        - ``tuple[pl.DataFrame, pl.DataFrame]``: inverse stateful context
          (reset=False, inverse=True, stateful=True).

        The first element is always the (transformed) input X; the second
        element, when present, is X_p.

    See Also
    --------
    - [`BaseActualTransformer`][yohou.base.transformer.BaseActualTransformer] : Base class for all transformers.
    - [`check_inputs`][yohou.utils.validation.check_inputs] : Low-level input validation helper.

    """
    if reset:
        # Fit context
        if X is None:
            raise ValueError("`X` cannot be None in fit context.")
        transformer_tags = getattr(transformer.__sklearn_tags__(), "transformer_tags", None)
        if getattr(transformer_tags, "accepts_irregular_grid", False):
            # The transformer declares it tolerates a non-uniform grid at fit (e.g. a
            # resampler that bins via group_by_dynamic), so the strict check is skipped
            # entirely rather than tried first. Trying it first would be unsound: on a
            # sub-day axis within its jitter tolerance the strict check does not raise,
            # it returns the median of the *unique* deltas, which a few outlier gaps
            # skew upward. Falling back only on ValueError would therefore use the
            # frequency-weighted (correct) median only when the skewed one failed
            # outright. On a uniform grid both agree, so the recorded interval is
            # unchanged there.
            validate_column_names(X)
            interval = representative_interval(X)
        else:
            interval = check_inputs(X, None)
        transformer.interval_ = interval
        transformer.feature_names_in_ = X.select(~cs.by_name("time")).columns
        transformer.n_features_in_ = len(transformer.feature_names_in_)
        transformer.X_schema_ = dict(X.select(~cs.by_name("time")).schema)
        return X

    # Transform/Inverse context (reset=False)
    if inverse:
        # Use X_t if provided, otherwise treat X as X_t (transformed data)
        if X_t is None:
            if X is None:
                raise ValueError("Either `X_t` or `X` must be provided for inverse transform.")
            X_t = X

        # Validate time columns
        check_time_column(X_t)
        if X_p is not None:
            check_time_column(X_p)

        if stateful and X_p is None:
            raise ValueError(
                "X_p cannot be None for stateful inverse transform. Provide the necessary previous untransformed data."
            )

        if observation_horizon is not None and observation_horizon > 0 and X_p is None:
            raise ValueError(
                "X_p cannot be None to invert a transform that has observation_horizon > 0. "
                "Provide the necessary previous untransformed data."
            )

        X_t_interval = None
        if len(X_t) >= 2:
            X_t_interval = check_interval_consistency(X_t)

        if X_p is not None and len(X_p) > 0 and observation_horizon is not None:
            if len(X_p) < observation_horizon:
                raise ValueError(
                    f"X_p must have at least {observation_horizon} rows (observation_horizon), "
                    f"but has only {len(X_p)} rows."
                )

            if len(X_p) > 1:
                X_p_interval = check_interval_consistency(X_p)
                if X_t_interval is not None and X_p_interval != X_t_interval:
                    raise ValueError(
                        f"Time intervals do not match: X_p has interval {X_p_interval}, "
                        f"but X_t has interval {X_t_interval}."
                    )

        return X_t, X_p

    # transform context
    if X is None:
        raise ValueError("`X` cannot be None for transform (when inverse=False).")
    check_time_column(X)
    X = check_schema(X, transformer.X_schema_)

    # A transformer that accepts an irregular grid (e.g. a resampler binning via
    # group_by_dynamic) skips the strict interval-consistency check at transform, the
    # same relaxation applied at fit above.
    transformer_tags = getattr(transformer.__sklearn_tags__(), "transformer_tags", None)
    accepts_irregular = getattr(transformer_tags, "accepts_irregular_grid", False)

    if not accepts_irregular and check_params.get("check_intervals", True) and len(X) >= 2:
        check_interval_consistency(X)

    # The accepts_irregular_grid gate below is unreachable today: the only transformer
    # that opts in (Downsampler) is stateless, so its _X_observed is empty and this
    # branch is skipped anyway. It is kept for a transformer that is both stateful and
    # irregular-tolerant (a time-based rolling window, say), which would otherwise
    # crash here, since the block resolves the interval with the strict check below.
    if (
        not accepts_irregular
        and check_params.get("check_continuity", True)
        and hasattr(transformer, "_X_observed")
        and len(transformer._X_observed) > 0
    ):
        interval = None
        if len(X) >= 2:
            interval = check_interval_consistency(X)
        check_continuity(
            transformer._X_observed,
            X,
            expected_interval=interval,
            check_intervals=(interval is not None),
        )

    return X