Skip to content

check_similarity_predict_matrix_shape

yohou.testing.check_similarity_predict_matrix_shape(similarity, y_calib, y_pred_calib)

Check predict returns an (n_pred, n_calib) weight matrix.

Parameters

Name Type Description Default
similarity BaseSimilarity

Similarity instance.

required
y_calib DataFrame

Calibration target frame.

required
y_pred_calib DataFrame

Calibration prediction frame.

required

Raises

Type Description
AssertionError

If predict does not return the expected (n_pred, n_calib) shape.

Source Code

Source code in src/yohou/testing/similarity.py
def check_similarity_predict_matrix_shape(similarity, y_calib: pl.DataFrame, y_pred_calib: pl.DataFrame) -> None:
    """Check ``predict`` returns an ``(n_pred, n_calib)`` weight matrix.

    Parameters
    ----------
    similarity : BaseSimilarity
        Similarity instance.
    y_calib : pl.DataFrame
        Calibration target frame.
    y_pred_calib : pl.DataFrame
        Calibration prediction frame.

    Raises
    ------
    AssertionError
        If ``predict`` does not return the expected ``(n_pred, n_calib)`` shape.

    """
    name = type(similarity).__name__
    sim = clone(similarity)
    sim.fit(y_calib, y_pred_calib)

    n_pred = 2
    query = y_pred_calib.head(n_pred)
    weights = sim.predict(query)
    assert isinstance(weights, np.ndarray), f"{name}.predict must return np.ndarray, got {type(weights)}"
    assert weights.shape == (n_pred, len(y_calib)), (
        f"{name}.predict returned shape {weights.shape}, expected ({n_pred}, {len(y_calib)})"
    )