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
|
||||
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
|
||||
|
||||
@@ -573,6 +589,31 @@ def test_data_preprocessor_get_required_columns_removes_duplicates():
|
||||
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():
|
||||
"""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})
|
||||
|
||||
Reference in New Issue
Block a user