72 lines
2.1 KiB
Python
72 lines
2.1 KiB
Python
"""Unit tests for TrainModelResult dataclass."""
|
|
|
|
import pandas as pd
|
|
import pytest
|
|
|
|
from model_manager.utils.models.train_model_params import TrainModelParams
|
|
from model_manager.utils.models.train_model_result import TrainModelResult
|
|
|
|
|
|
@pytest.fixture
|
|
def sample_params() -> TrainModelParams:
|
|
"""Minimal TrainModelParams for TrainModelResult tests."""
|
|
return TrainModelParams.from_dict(
|
|
{
|
|
'variable_columns': ['a'],
|
|
'target_variable': 't',
|
|
'bucket_name': 'b',
|
|
'file_name': 'f.csv',
|
|
'line_separator': '\n',
|
|
'decimal_separator': '.',
|
|
'date_column': 'timestamp',
|
|
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
|
'train_size': 80,
|
|
'shuffle': True,
|
|
'random_state': 42,
|
|
'experiment_run_id': 1,
|
|
'model_name': 'Linear Regression',
|
|
'val_file_name': None,
|
|
'data_model_kwargs': {},
|
|
'model_kwargs': {},
|
|
'opt_params': {},
|
|
'model_type': 'linear_regression',
|
|
'model_id': None,
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def sample_frames():
|
|
train = pd.DataFrame({'a': [1, 2], 't': [1.0, 2.0]})
|
|
val = pd.DataFrame({'a': [3], 't': [3.0]})
|
|
return train, val
|
|
|
|
|
|
def test_train_model_result_creation(sample_params, sample_frames):
|
|
train, val = sample_frames
|
|
result = TrainModelResult(params=sample_params, train_data=train, val_data=val)
|
|
assert result.params is sample_params
|
|
assert result.train_data.equals(train)
|
|
assert result.val_data.equals(val)
|
|
assert result.run_name is None
|
|
|
|
|
|
def test_train_model_result_optional_paths(sample_params, sample_frames):
|
|
train, val = sample_frames
|
|
result = TrainModelResult(
|
|
params=sample_params,
|
|
train_data=train,
|
|
val_data=val,
|
|
run_name='run-1',
|
|
run_id='rid',
|
|
run_dir='/tmp/x',
|
|
mse_val=0.1,
|
|
mae_val=0.2,
|
|
r2_val=0.99,
|
|
)
|
|
assert result.run_name == 'run-1'
|
|
assert result.run_id == 'rid'
|
|
assert result.run_dir == '/tmp/x'
|
|
assert result.mse_val == 0.1
|