Skip to content

Commit 553463a

Browse files
committed
Merge branch 'feature/model-selector' into develop
2 parents 43c663b + 75b7cd4 commit 553463a

2 files changed

Lines changed: 234 additions & 6 deletions

File tree

ats/tests/test_utils.py

Lines changed: 146 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -156,3 +156,149 @@ def test_plotter_from_df_missing_columns_raises(self):
156156
df_missing = self.df.drop(columns=["fitness"])
157157
with self.assertRaises(KeyError):
158158
plot_from_df(df_missing, "avg_err", "fitness", {})
159+
160+
def test_model_selector(self):
161+
from ..utils import select_models
162+
models = select_models()
163+
self.assertTrue(len(models) > 0)
164+
165+
# Use MinMaxAnomalyDetector as a reference for capability-based selection
166+
from ..anomaly_detectors.naive.minmax import MinMaxAnomalyDetector
167+
requested_capabilities = {
168+
"training": {
169+
"mode": None,
170+
"update": False
171+
},
172+
"inference": {
173+
"streaming": False,
174+
"dependency": "series",
175+
"granularity": "point-labels"
176+
},
177+
"data": {
178+
"dimensionality": "univariate-single",
179+
"sampling": "irregular"
180+
}
181+
}
182+
183+
selected = select_models(requested_capabilities)
184+
self.assertIn(MinMaxAnomalyDetector, selected)
185+
186+
def test_model_selector_no_filter_returns_all(self):
187+
from ..utils import select_models
188+
from ..anomaly_detectors import (
189+
MinMaxAnomalyDetector, ZScoreAnomalyDetector,
190+
COMAnomalyDetector, HARAnomalyDetector, NHARAnomalyDetector,
191+
IFSOMAnomalyDetector, LinearRegressionAnomalyDetector, LSTMAnomalyDetector,
192+
)
193+
all_expected = {
194+
MinMaxAnomalyDetector, ZScoreAnomalyDetector,
195+
COMAnomalyDetector, HARAnomalyDetector, NHARAnomalyDetector,
196+
IFSOMAnomalyDetector, LinearRegressionAnomalyDetector, LSTMAnomalyDetector,
197+
}
198+
models = select_models()
199+
self.assertEqual(set(models), all_expected)
200+
201+
def test_model_selector_by_training_mode_none(self):
202+
from ..utils import select_models
203+
from ..anomaly_detectors import MinMaxAnomalyDetector, ZScoreAnomalyDetector
204+
selected = select_models({"training": {"mode": None}})
205+
self.assertEqual(set(selected), {MinMaxAnomalyDetector, ZScoreAnomalyDetector})
206+
207+
def test_model_selector_by_training_mode_unsupervised(self):
208+
from ..utils import select_models
209+
from ..anomaly_detectors import (
210+
COMAnomalyDetector, HARAnomalyDetector, NHARAnomalyDetector,
211+
IFSOMAnomalyDetector,
212+
)
213+
selected = select_models({"training": {"mode": "unsupervised"}})
214+
self.assertEqual(set(selected), {
215+
COMAnomalyDetector, HARAnomalyDetector, NHARAnomalyDetector,
216+
IFSOMAnomalyDetector,
217+
})
218+
219+
def test_model_selector_by_training_mode_semi_supervised(self):
220+
from ..utils import select_models
221+
from ..anomaly_detectors import LinearRegressionAnomalyDetector, LSTMAnomalyDetector
222+
selected = select_models({"training": {"mode": "semi-supervised"}})
223+
self.assertEqual(set(selected), {LinearRegressionAnomalyDetector, LSTMAnomalyDetector})
224+
225+
def test_model_selector_by_inference_dependency_window(self):
226+
from ..utils import select_models
227+
from ..anomaly_detectors import LinearRegressionAnomalyDetector, LSTMAnomalyDetector
228+
selected = select_models({"inference": {"dependency": "window"}})
229+
self.assertEqual(set(selected), {LinearRegressionAnomalyDetector, LSTMAnomalyDetector})
230+
231+
def test_model_selector_by_inference_granularity_series(self):
232+
from ..utils import select_models
233+
from ..anomaly_detectors import IFSOMAnomalyDetector
234+
selected = select_models({"inference": {"granularity": "series"}})
235+
self.assertEqual(set(selected), {IFSOMAnomalyDetector})
236+
237+
def test_model_selector_by_sampling_irregular(self):
238+
from ..utils import select_models
239+
from ..anomaly_detectors import MinMaxAnomalyDetector, ZScoreAnomalyDetector
240+
selected = select_models({"data": {"sampling": "irregular"}})
241+
self.assertEqual(set(selected), {MinMaxAnomalyDetector, ZScoreAnomalyDetector})
242+
243+
def test_model_selector_by_dimensionality_scalar(self):
244+
"""Requesting a single dimensionality value matches models that include it in their list."""
245+
from ..utils import select_models
246+
from ..anomaly_detectors import (
247+
MinMaxAnomalyDetector, COMAnomalyDetector, HARAnomalyDetector,
248+
NHARAnomalyDetector, IFSOMAnomalyDetector,
249+
)
250+
selected = select_models({"data": {"dimensionality": "univariate-multi"}})
251+
self.assertEqual(set(selected), {
252+
MinMaxAnomalyDetector, COMAnomalyDetector, HARAnomalyDetector,
253+
NHARAnomalyDetector, IFSOMAnomalyDetector,
254+
})
255+
256+
def test_model_selector_by_dimensionality_list(self):
257+
"""Requesting a list of dimensionalities matches models supporting ALL of them (subset check)."""
258+
from ..utils import select_models
259+
from ..anomaly_detectors import (
260+
MinMaxAnomalyDetector, COMAnomalyDetector, HARAnomalyDetector,
261+
NHARAnomalyDetector,
262+
)
263+
selected = select_models({"data": {"dimensionality": ["univariate-multi", "multivariate-single"]}})
264+
self.assertEqual(set(selected), {
265+
MinMaxAnomalyDetector, COMAnomalyDetector, HARAnomalyDetector,
266+
NHARAnomalyDetector,
267+
})
268+
269+
def test_model_selector_multiple_sections(self):
270+
"""Filtering across multiple capability sections at once."""
271+
from ..utils import select_models
272+
from ..anomaly_detectors import MinMaxAnomalyDetector, ZScoreAnomalyDetector
273+
selected = select_models({
274+
"training": {"mode": None},
275+
"data": {"sampling": "irregular"},
276+
})
277+
self.assertEqual(set(selected), {MinMaxAnomalyDetector, ZScoreAnomalyDetector})
278+
279+
def test_model_selector_multiple_fields_narrow(self):
280+
"""Combining multiple fields to narrow down to a unique model."""
281+
from ..utils import select_models
282+
from ..anomaly_detectors import IFSOMAnomalyDetector
283+
selected = select_models({
284+
"training": {"mode": "unsupervised"},
285+
"inference": {"granularity": "series"},
286+
})
287+
self.assertEqual(set(selected), {IFSOMAnomalyDetector})
288+
289+
def test_model_selector_no_match(self):
290+
"""Requesting a combination no model satisfies returns an empty list."""
291+
from ..utils import select_models
292+
selected = select_models({
293+
"training": {"mode": "supervised"},
294+
})
295+
self.assertEqual(selected, [])
296+
297+
def test_model_selector_no_match_contradictory(self):
298+
"""Contradictory constraints across sections return empty."""
299+
from ..utils import select_models
300+
selected = select_models({
301+
"training": {"mode": None},
302+
"data": {"sampling": "regular"},
303+
})
304+
self.assertEqual(selected, [])

ats/utils.py

Lines changed: 88 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ def generate_timeseries_df(start='2025-06-10 14:00:00', tz='UTC', freq='h', ent
3434
df.index.name = 'timestamp'
3535
return df
3636

37+
3738
def convert_timeseries_df_to_timeseries(timeseries_df):
3839
timeseries = TimeSeries.from_df(timeseries_df)
3940

@@ -44,6 +45,7 @@ def convert_timeseries_df_to_timeseries(timeseries_df):
4445

4546
return timeseries
4647

48+
4749
def convert_timeseries_to_timeseries_df(timeseries):
4850
timeseries_df = TimeSeries.to_df(timeseries)
4951

@@ -58,6 +60,7 @@ def convert_timeseries_to_timeseries_df(timeseries):
5860

5961
return timeseries_df
6062

63+
6164
def plot_timeseries_df(timeseries_df, *args, **kwargs):
6265

6366
timeseries = convert_timeseries_df_to_timeseries(timeseries_df)
@@ -84,7 +87,7 @@ def normalize_parameter(df, parameter):
8487
return pd.Series(1.0, index=df.index, name=parameter)
8588
else:
8689
return (df[parameter] - min_parameter) / (max_parameter - min_parameter).astype(float)
87-
90+
8891
def normalize_df(df, parameters_subset=None,save=False):
8992
"""
9093
Normalizes a DataFrame using (value-min)/(max-min).
@@ -173,7 +176,8 @@ def plot_3d_interactive(df,x="avg_err",y="max_err",z="ks_pvalue",color="fitness"
173176
return None
174177
except Exception as e:
175178
logger.error("Unexpected error while creating 3D interactive plot: %s", e)
176-
return None
179+
return None
180+
177181

178182
def save_df_to_csv(df, outputfile="output.csv"):
179183
"""
@@ -271,7 +275,7 @@ def find_best_parameter(df, parameter, mode="min"):
271275
except Exception as e:
272276
logger.error(f"Error finding {mode} for '{parameter}': {e} ({type(e).__name__})")
273277
return None
274-
278+
275279

276280
def plot_from_df(df, x,y,fixed_parameters=None):
277281
"""
@@ -296,16 +300,16 @@ def plot_from_df(df, x,y,fixed_parameters=None):
296300
if key not in df_filtered.columns:
297301
logger.warning(f"'{key}' not in DataFrame columns. Skipping filter.")
298302
continue
299-
303+
300304
if isinstance(val, (list, tuple, set)):
301305
df_filtered = df_filtered[df_filtered[key].isin(val)]
302306
else:
303307
df_filtered = df_filtered[df_filtered[key] == val]
304-
308+
305309
df_filtered = df_filtered.sort_values(by=x)
306310

307311
context_info = " | ".join(f"{k}={v}" for k, v in (fixed_parameters or {}).items())
308-
312+
309313
try:
310314
fig, ax = plt.subplots(figsize=(8, 5))
311315
ax.plot(df_filtered[x], df_filtered[y], marker="o")
@@ -432,6 +436,7 @@ def timeseries_df_to_list_of_timeseries_df(timeseries_df, anomaly_labels=False):
432436
def list_of_timeseries_df_to_timeseries_df(list_of_timeseries_df):
433437
return pd.concat(list_of_timeseries_df, axis=1)
434438

439+
435440
def ensure_full_reproducibility(seed=0):
436441
random.seed(seed)
437442
np.random.seed(seed)
@@ -443,3 +448,80 @@ def ensure_full_reproducibility(seed=0):
443448
tf.random.set_seed(seed)
444449
tf.config.experimental.enable_op_determinism()
445450
os.environ["TF_DETERMINISTIC_OPS"] = "1"
451+
452+
453+
def select_models(capabilities={}):
454+
455+
import inspect
456+
457+
def get_classes(module):
458+
"""
459+
Return a list of class objects defined/imported in `module`.
460+
"""
461+
return [
462+
obj
463+
for obj in vars(module).values()
464+
if inspect.isclass(obj)
465+
]
466+
467+
from . import anomaly_detectors
468+
model_classes = get_classes(anomaly_detectors)
469+
logger.debug('Searching in #{} models'.format(len(model_classes)))
470+
471+
selected_model_classes = []
472+
473+
def _matches_capability(actual_value, requested_value):
474+
"""
475+
Return True if the requested_value is compatible with actual_value.
476+
Handles scalar vs list-like values on either side.
477+
"""
478+
actual_is_list = isinstance(actual_value, (list, set, tuple))
479+
requested_is_list = isinstance(requested_value, (list, set, tuple))
480+
481+
if requested_is_list:
482+
requested_set = set(requested_value)
483+
if actual_is_list:
484+
return requested_set.issubset(set(actual_value))
485+
return actual_value in requested_set
486+
487+
# requested is scalar
488+
if actual_is_list:
489+
return requested_value in actual_value
490+
return actual_value == requested_value
491+
492+
for model_class in model_classes:
493+
logger.debug('Checking capabilities against #{}'.format(model_class.__name__))
494+
495+
# If capabilities match, add them to the selected_model_classes list
496+
if not capabilities:
497+
selected_model_classes.append(model_class)
498+
continue
499+
500+
try:
501+
model_capabilities = model_class.capabilities
502+
except Exception as e:
503+
logger.warning(
504+
f'Skipping {model_class.__name__} due to invalid capabilities: {e}'
505+
)
506+
continue
507+
508+
matches = True
509+
for section, fields in capabilities.items():
510+
if section not in model_capabilities:
511+
matches = False
512+
break
513+
for field, requested_value in fields.items():
514+
if field not in model_capabilities[section]:
515+
matches = False
516+
break
517+
actual_value = model_capabilities[section][field]
518+
if not _matches_capability(actual_value, requested_value):
519+
matches = False
520+
break
521+
if not matches:
522+
break
523+
524+
if matches:
525+
selected_model_classes.append(model_class)
526+
527+
return selected_model_classes

0 commit comments

Comments
 (0)