Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 8 additions & 9 deletions aeon/classification/dictionary_based/_redcomets.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,7 +160,7 @@ def _fit(self, X, y):
self.sax_clfs,
) = self._build_univariate_ensemble(X_concat, y)

elif self.variant in [4, 5, 6, 7, 8, 9]: # Ensemble
else: # Ensemble (variants 4-9)
(
self.sfa_transforms,
self.sfa_clfs,
Expand Down Expand Up @@ -396,8 +396,8 @@ def _predict_proba(self, X) -> np.ndarray:
if self.variant in [1, 2, 3]: # Concatenate
X_concat = X.reshape(*X.shape[:-2], -1)
return self._predict_proba_unvivariate(X_concat)
elif self.variant in [4, 5, 6, 7, 8, 9]:
return self._predict_proba_dimension_ensemble(X) # Ensemble
else: # Ensemble (variants 4-9)
return self._predict_proba_dimension_ensemble(X)

def _predict_proba_unvivariate(self, X) -> np.ndarray:
"""Predicts labels probabilities for sequences in univariate X.
Expand Down Expand Up @@ -473,8 +473,7 @@ def _predict_proba_dimension_ensemble(self, X) -> np.ndarray:
if self.variant in [6, 7, 8, 9]:
dimension_pred_mats = None
for sfa, (rf, _) in zip(sfa_transforms, sfa_clfs):
sfa_dics = sfa.transform_words(X_d)
X_sfa = sfa_dics[:, 0, :]
X_sfa = sfa.transform_words(X_d)[0]

rf_pred_mat = rf.predict_proba(X_sfa)

Expand All @@ -486,7 +485,7 @@ def _predict_proba_dimension_ensemble(self, X) -> np.ndarray:
(ensemble_pred_mats, [rf_pred_mat])
)

elif self.variant in [6, 7, 8, 9]:
else: # variants 6-9
if dimension_pred_mats is None:
dimension_pred_mats = [rf_pred_mat]
else:
Expand All @@ -507,7 +506,7 @@ def _predict_proba_dimension_ensemble(self, X) -> np.ndarray:
(ensemble_pred_mats, [rf_pred_mat])
)

elif self.variant in [6, 7, 8, 9]:
else: # variants 6-9
if dimension_pred_mats is None:
dimension_pred_mats = [rf_pred_mat]
else:
Expand All @@ -518,7 +517,7 @@ def _predict_proba_dimension_ensemble(self, X) -> np.ndarray:
if self.variant in [6, 7, 8, 9]:
if self.variant in [6, 7]:
fused_dimension_pred_mat = np.sum(dimension_pred_mats, axis=0)
elif self.variant in [8, 9]:
else: # variants 8, 9
weights = np.array(
[np.mean(mat.max(axis=1)) for mat in dimension_pred_mats]
).reshape(-1, 1)
Expand All @@ -535,7 +534,7 @@ def _predict_proba_dimension_ensemble(self, X) -> np.ndarray:

if self.variant in [4, 6, 7]:
pred_mat = np.sum(np.array(ensemble_pred_mats), axis=0)
elif self.variant in [5, 8, 9]:
else: # variants 5, 8, 9
weights = np.array(
[np.mean(mat.max(axis=1)) for mat in ensemble_pred_mats]
).reshape(-1, 1)
Expand Down
260 changes: 155 additions & 105 deletions aeon/classification/dictionary_based/tests/test_redcomets.py
Original file line number Diff line number Diff line change
@@ -1,107 +1,157 @@
"""REDCOMETS test code."""

__maintainer__ = []

# from sys import platform
#
# import numpy as np
# import pytest
#
# from aeon.classification.dictionary_based import REDCOMETS
# from aeon.datasets import load_basic_motions, load_unit_test
# from aeon.utils.validation._dependencies import _check_soft_dependencies
#
#
# @pytest.mark.skipif(
# not _check_soft_dependencies(
# "imbalanced-learn",
# package_import_alias={"imbalanced-learn": "imblearn"},
# severity="none",
# ),
# reason="skip test if required soft dependency imbalanced-learn not available",
# )
# def test_redcomets_score_univariate():
# """Test of REDCOMETS train estimate on unit test data."""
# # load unit test data
# X_train, y_train = load_unit_test(split="train")
# X_test, y_test = load_unit_test(split="test")
#
# def test_variant(v, expected_result):
# # train REDCOMETS-<v>
# redcomets = REDCOMETS(variant=v, n_trees=3, random_state=0)
# redcomets.fit(X_train, y_train)
#
# score = redcomets.score(X_test, y_test)
#
# assert isinstance(score, float)
#
# # We cannot guarantee same results on ARM macOS
# if platform != "darwin":
# np.testing.assert_almost_equal(score, expected_result, decimal=4)
#
# test_variant(1, 0.7272)
# test_variant(2, 0.6818)
# test_variant(3, 0.7272)
#
#
# @pytest.mark.skipif(
# not _check_soft_dependencies(
# "imbalanced-learn",
# package_import_alias={"imbalanced-learn": "imblearn"},
# severity="none",
# ),
# reason="skip test if required soft dependency imbalanced-learn not available",
# )
# def test_redcomets_score_multivariate():
# """Test of REDCOMETS train estimate on unit test data."""
# # load unit test data
# X_train, y_train = load_basic_motions(split="train")
# X_test, y_test = load_basic_motions(split="test")
#
# def test_variant(v, expected_result):
# # train REDCOMETS-<v>
# redcomets = REDCOMETS(variant=v, n_trees=3, random_state=0)
# redcomets.fit(X_train, y_train)
#
# score = redcomets.score(X_test, y_test)
#
# assert isinstance(score, float)
#
# # We cannot guarantee same results on ARM macOS
# if platform != "darwin":
# np.testing.assert_almost_equal(score, expected_result, decimal=4)
#
# test_variant(1, 0.95)
# test_variant(2, 0.975)
# test_variant(3, 0.975)
# test_variant(4, 0.875)
# test_variant(5, 0.875)
# test_variant(6, 0.875)
# test_variant(7, 0.875)
# test_variant(8, 0.875)
# test_variant(9, 0.875)
#
#
# @pytest.mark.skipif(
# not _check_soft_dependencies(
# "imbalanced-learn",
# package_import_alias={"imbalanced-learn": "imblearn"},
# severity="none",
# ),
# reason="skip test if required soft dependency imbalanced-learn not available",
# )
# def test_redcomets_lens_generation():
# """Test of REDCOMETS random lens generation."""
# # load unit test data
# X, y = load_unit_test()
#
# # Generate 10 random lenses
# redcomets = REDCOMETS(random_state=0)
# lenses = redcomets._get_random_lenses(np.squeeze(X), 10)
#
# assert len(lenses) == 10
# assert isinstance(lenses, list)
#
# for w, a in lenses:
# assert isinstance(w, int)
# assert isinstance(a, int)
import numpy as np
import pytest
from sklearn.utils import check_random_state

from aeon.classification.dictionary_based import REDCOMETS

N_PER_CLASS = 10
N_TIMEPOINTS = 48 # long enough to yield >=2 SFA and >=2 SAX lenses per view
N_CLASSES = 2


def _labelled_panel(n_channels, n_per_class=N_PER_CLASS, random_state=0):
"""Return a balanced random panel with ``N_CLASSES`` classes.

Values are random: the tests assert output structure (shape, valid labels,
normalised probabilities), not classification accuracy.
"""
rng = check_random_state(random_state)
n_cases = n_per_class * N_CLASSES
X = rng.standard_normal((n_cases, n_channels, N_TIMEPOINTS))
y = np.repeat(np.arange(N_CLASSES), n_per_class)
return X, y


def _assert_valid_output(clf, X):
"""Check output structure: normalised probabilities and in-vocabulary labels."""
proba = clf.predict_proba(X)
pred = clf.predict(X)

assert proba.shape == (X.shape[0], clf.n_classes_)
np.testing.assert_allclose(proba.sum(axis=1), 1.0)
assert pred.shape == (X.shape[0],)
assert set(pred).issubset(set(clf.classes_))


@pytest.mark.parametrize("variant", [1, 2, 3])
def test_redcomets_univariate_variants(variant):
"""Univariate variants 1-3 fit and produce well-formed predictions."""
X, y = _labelled_panel(n_channels=1)
clf = REDCOMETS(variant=variant, n_trees=3, random_state=0)
clf.fit(X, y)
_assert_valid_output(clf, X)


@pytest.mark.parametrize("variant", [1, 2, 3])
def test_redcomets_multivariate_concatenate_variants(variant):
"""Variants 1-3 handle multivariate input by concatenating channels."""
X, y = _labelled_panel(n_channels=3)
clf = REDCOMETS(variant=variant, n_trees=3, random_state=0)
clf.fit(X, y)
_assert_valid_output(clf, X)


@pytest.mark.parametrize("variant", [4, 5, 6, 7, 8, 9])
def test_redcomets_dimension_ensemble_variants(variant):
"""Variants 4-9 build and fuse a per-channel ensemble on multivariate input.

These variants exercise the dimension-ensemble build and the variant-specific
fusion in ``_predict_proba_dimension_ensemble`` (plain sum vs. confidence
weighting, at both the per-channel and cross-channel stages).
"""
X, y = _labelled_panel(n_channels=3)
clf = REDCOMETS(variant=variant, n_trees=3, random_state=0)
clf.fit(X, y)
_assert_valid_output(clf, X)


def test_redcomets_deterministic():
"""A fixed random_state gives identical predictions across fits."""
X, y = _labelled_panel(n_channels=3)

pred1 = REDCOMETS(variant=5, n_trees=3, random_state=0).fit(X, y).predict(X)
pred2 = REDCOMETS(variant=5, n_trees=3, random_state=0).fit(X, y).predict(X)

np.testing.assert_array_equal(pred1, pred2)


def test_redcomets_balanced_input_needs_no_oversampling():
"""Already-balanced classes fit without invoking the oversampling branch."""
X, y = _labelled_panel(n_channels=1) # N_PER_CLASS each, balanced
assert np.unique(y, return_counts=True)[1].tolist() == [N_PER_CLASS, N_PER_CLASS]

clf = REDCOMETS(variant=1, n_trees=3, random_state=0)
clf.fit(X, y)
_assert_valid_output(clf, X)


def test_redcomets_imbalanced_input_uses_smote():
"""An imbalanced class large enough for neighbour search is SMOTE-oversampled.

The minority class has more than five samples, exercising the capped
neighbour-count SMOTE path rather than the fallback.
"""
X, _ = _labelled_panel(n_channels=1, n_per_class=14)
y = np.array([0] * 20 + [1] * 8) # minority > 5 -> capped SMOTE neighbours

clf = REDCOMETS(variant=1, n_trees=3, random_state=0)
clf.fit(X, y)
assert set(clf.classes_) == {0, 1}
_assert_valid_output(clf, X)


def test_redcomets_tiny_minority_uses_random_oversampler():
"""A minority class too small for SMOTE falls back to random oversampling.

With two minority samples the SMOTE neighbour count drops below one, so
REDCOMETS must fall back to RandomOverSampler and still fit on both classes.
"""
X, _ = _labelled_panel(n_channels=1, n_per_class=10)
y = np.array([0] * 18 + [1] * 2)

clf = REDCOMETS(variant=1, n_trees=3, random_state=0)
clf.fit(X, y)
assert set(clf.classes_) == {0, 1}
_assert_valid_output(clf, X)


@pytest.mark.parametrize("bad_variant", [0, 10])
def test_redcomets_rejects_invalid_variant(bad_variant):
"""Variants outside 1-9 are rejected at construction."""
with pytest.raises(AssertionError):
REDCOMETS(variant=bad_variant)


@pytest.mark.parametrize("bad_perc", [0, 101])
def test_redcomets_rejects_invalid_perc_length(bad_perc):
"""perc_length must lie in (0, 100]."""
with pytest.raises(AssertionError):
REDCOMETS(perc_length=bad_perc)


def test_redcomets_univariate_rejects_ensemble_variant():
"""Dimension-ensemble variants 4-9 require multivariate input."""
X, y = _labelled_panel(n_channels=1)
clf = REDCOMETS(variant=4, n_trees=3, random_state=0)
with pytest.raises(AssertionError):
clf.fit(X, y)


def test_redcomets_test_params_are_valid():
"""The documented test parameters construct a valid REDCOMETS instance."""
params = REDCOMETS._get_test_params()
assert params["variant"] in range(1, 10)
REDCOMETS(**params) # construction asserts pass


def test_redcomets_declares_no_imbalanced_learn_dependency():
"""REDCOMETS no longer depends on imbalanced-learn (gh-3654)."""
deps = REDCOMETS(random_state=0).get_tag("python_dependencies", None)
if deps is None:
deps = []
if isinstance(deps, str):
deps = [deps]
assert "imblearn" not in deps
assert "imbalanced-learn" not in deps
Loading
Loading