SIENTIAPDE-1255: Add tests to ensure get_required_columns removes columns from self_operations and cross_operations.
This commit is contained in:
@@ -6,7 +6,23 @@ import numpy as np
|
|||||||
import pandas as pd
|
import pandas as pd
|
||||||
from pytest import raises
|
from pytest import raises
|
||||||
|
|
||||||
from model_manager.sientia.models import DataPreprocessor, LinearRegressionModel
|
from model_manager.sientia.models import (
|
||||||
|
DataPreprocessor,
|
||||||
|
LinearRegressionModel,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class _IterableWithContains:
|
||||||
|
def __init__(self, iterable, contains_values):
|
||||||
|
self._iterable = iterable
|
||||||
|
self._contains = set(contains_values)
|
||||||
|
|
||||||
|
def __iter__(self):
|
||||||
|
return iter(self._iterable)
|
||||||
|
|
||||||
|
def __contains__(self, item):
|
||||||
|
return item in self._contains
|
||||||
|
|
||||||
|
|
||||||
# LinearRegressionModel Tests
|
# LinearRegressionModel Tests
|
||||||
|
|
||||||
@@ -573,6 +589,31 @@ def test_data_preprocessor_get_required_columns_removes_duplicates():
|
|||||||
assert 'var1' not in result or result.count('var1') <= 1
|
assert 'var1' not in result or result.count('var1') <= 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_data_preprocessor_get_required_columns_removes_self_operations_branch():
|
||||||
|
"""Ensure line 278 removes columns present in self_operations iterable."""
|
||||||
|
preprocessor = DataPreprocessor(
|
||||||
|
self_operations=_IterableWithContains(['{var1}_{pow}_{2}'], contains_values=['var1'])
|
||||||
|
)
|
||||||
|
existing_columns: list[str] = []
|
||||||
|
|
||||||
|
result = preprocessor.get_required_columns(existing_columns)
|
||||||
|
|
||||||
|
assert 'var1' not in result
|
||||||
|
|
||||||
|
|
||||||
|
def test_data_preprocessor_get_required_columns_removes_cross_operations_branch():
|
||||||
|
"""Ensure line 285 removes columns present in cross_operations iterable."""
|
||||||
|
preprocessor = DataPreprocessor(
|
||||||
|
cross_operations=_IterableWithContains(['{var1}_{*}_{var2}'], contains_values=['var1'])
|
||||||
|
)
|
||||||
|
existing_columns: list[str] = []
|
||||||
|
|
||||||
|
result = preprocessor.get_required_columns(existing_columns)
|
||||||
|
|
||||||
|
assert 'var1' not in result
|
||||||
|
assert 'var2' in result
|
||||||
|
|
||||||
|
|
||||||
def test_data_preprocessor_get_required_columns_removes_created_lags():
|
def test_data_preprocessor_get_required_columns_removes_created_lags():
|
||||||
"""Test get_required_columns removes columns from created_lags when column is in created_lags dict - covers line 292."""
|
"""Test get_required_columns removes columns from created_lags when column is in created_lags dict - covers line 292."""
|
||||||
preprocessor = DataPreprocessor(created_lags={'var1': 1, 'var2': 1})
|
preprocessor = DataPreprocessor(created_lags={'var1': 1, 'var2': 1})
|
||||||
|
|||||||
Reference in New Issue
Block a user