MYTHOSLIVE / tests /test_models_with_features.py
3VVM's picture
Upload MYTHOSLIVE project from A2 ZIP
9831ced verified
Raw History Blame Contribute Delete
7.82 kB
import importlib.util
import unittest
import numpy as np
import pandas as pd
from features.feature_pipeline import compute_feature_frame
def _have(module_name: str) -> bool:
return importlib.util.find_spec(module_name) is not None
_HAVE_STATSMODELS = _have('statsmodels')
_HAVE_ARCH = _have('arch')
_HAVE_MOIRAI = _have('uni2ts') and _have('gluonts')
_HAVE_TIMESFM = _have('timesfm') and _have('torch')
def _synthetic_ohlcv(n: int, seed: int=42) -> pd.DataFrame:
rng = np.random.default_rng(seed)
steps = rng.normal(loc=0.0, scale=0.4, size=n)
close = 100 + np.cumsum(steps)
close = np.abs(close) + 5.0
open_ = close + rng.normal(0, 0.1, n)
high = np.maximum(open_, close) + np.abs(rng.normal(0, 0.2, n))
low = np.minimum(open_, close) - np.abs(rng.normal(0, 0.2, n))
volume = rng.integers(100, 10000, n).astype(float)
idx = pd.date_range('2024-01-01', periods=n, freq='min')
return pd.DataFrame({'Open': open_, 'High': high, 'Low': low, 'Close': close, 'Volume': volume}, index=idx)
def _warm_close_and_features(df: pd.DataFrame, feats: pd.DataFrame):
complete = ~feats.isna().any(axis=1).to_numpy()
start = int(np.argmax(complete)) if complete.any() else len(df)
return (df['Close'].reset_index(drop=True).iloc[start:], feats.iloc[start:].reset_index(drop=True))
class TestArimaFamilyAlignment(unittest.TestCase):
def test_returns_alignment_has_no_off_by_one(self):
n = 10
prices = pd.Series(100 + np.arange(n, dtype=float))
features = pd.Series([float(i) for i in range(n)])
log_returns = np.log(prices / prices.shift(1)).dropna().reset_index(drop=True)
exog = features.iloc[1:].reset_index(drop=True)
self.assertEqual(len(exog), len(log_returns))
for k in range(len(log_returns)):
self.assertEqual(exog[k], float(k + 1))
@unittest.skipUnless(_HAVE_STATSMODELS, 'statsmodels not installed in this environment')
class TestArimaWithFeatures(unittest.TestCase):
def setUp(self):
self.df = _synthetic_ohlcv(120)
self.feats = compute_feature_frame(self.df)
self.close, self.warm_feats = _warm_close_and_features(self.df, self.feats)
def test_baseline_unchanged_shape(self):
from models.arima_model import ArimaModel
m = ArimaModel(max_p=2, max_d=1, max_q=2)
out = m.predict(self.close, horizon=3)
self.assertEqual(len(out), 3)
self.assertTrue(all(np.isfinite(out)))
def test_with_features_runs_and_is_finite(self):
from models.arima_model import ArimaModel
m = ArimaModel(max_p=2, max_d=1, max_q=2)
out = m.predict(self.close, horizon=3, features=self.warm_feats)
self.assertEqual(len(out), 3)
self.assertTrue(all(np.isfinite(out)))
def test_mismatched_features_length_raises(self):
from models.arima_model import ArimaModel
m = ArimaModel(max_p=2, max_d=1, max_q=2)
short_window = self.warm_feats.iloc[:-5].reset_index(drop=True)
with self.assertRaises(ValueError):
m.predict(self.close, horizon=3, features=short_window)
@unittest.skipUnless(_HAVE_STATSMODELS, 'statsmodels not installed in this environment')
class TestAutoArimaWithFeatures(unittest.TestCase):
def setUp(self):
self.df = _synthetic_ohlcv(120)
self.feats = compute_feature_frame(self.df)
self.close, self.warm_feats = _warm_close_and_features(self.df, self.feats)
def test_with_features_runs_and_is_finite(self):
from models.auto_arima_model import AutoArimaModel
m = AutoArimaModel(max_p=3, max_q=3, max_d=1)
out = m.predict(self.close, horizon=2, features=self.warm_feats)
self.assertEqual(len(out), 2)
self.assertTrue(all(np.isfinite(out)))
@unittest.skipUnless(_HAVE_STATSMODELS and _HAVE_ARCH, 'statsmodels+arch not installed in this environment')
class TestArimaGarchWithFeatures(unittest.TestCase):
def setUp(self):
self.df = _synthetic_ohlcv(150)
self.feats = compute_feature_frame(self.df)
self.close, self.warm_feats = _warm_close_and_features(self.df, self.feats)
def test_with_features_runs_and_is_finite(self):
from models.arima_garch_model import ArimaGarchModel
m = ArimaGarchModel(max_p=2, max_q=2)
out = m.predict(self.close, horizon=2, features=self.warm_feats)
self.assertEqual(len(out), 2)
self.assertTrue(all(np.isfinite(out)))
self.assertIsNotNone(m.last_volatility_forecast)
@unittest.skipUnless(_HAVE_MOIRAI, 'uni2ts/gluonts not installed in this environment')
class TestMoiraiWithFeatures(unittest.TestCase):
def test_with_features_runs_and_is_finite(self):
from models.moirai_model import MoiraiModel
df = _synthetic_ohlcv(60)
feats = compute_feature_frame(df)
close, warm_feats = _warm_close_and_features(df, feats)
m = MoiraiModel(context_length=60, num_samples=10)
out = m.predict(close, horizon=2, features=warm_feats)
self.assertEqual(len(out), 2)
self.assertTrue(all(np.isfinite(out)))
@unittest.skipUnless(_HAVE_TIMESFM, 'timesfm/torch not installed in this environment')
class TestTimesFMWithFeatures(unittest.TestCase):
def test_with_features_runs_and_is_finite(self):
from models.timesfm_model import TimesFMModel
df = _synthetic_ohlcv(64)
feats = compute_feature_frame(df)
close, warm_feats = _warm_close_and_features(df, feats)
m = TimesFMModel(max_context=64)
out = m.predict(close, horizon=2, features=warm_feats)
self.assertEqual(len(out), 2)
self.assertTrue(all(np.isfinite(out)))
def test_short_input_below_compiled_context_is_finite(self):
from models.timesfm_model import TimesFMModel
df = _synthetic_ohlcv(60, seed=99)
close = 40000 + df['Close'].reset_index(drop=True)
m = TimesFMModel(max_context=512)
out = m.predict(close, horizon=2)
self.assertEqual(len(out), 2)
self.assertTrue(all(np.isfinite(out)), f'TimesFM NaN on short window: {out}')
def test_repeated_window_lengths_stay_finite(self):
from models.timesfm_model import TimesFMModel
df = _synthetic_ohlcv(120, seed=7)
close = 40000 + df['Close'].reset_index(drop=True)
m = TimesFMModel(max_context=512)
for w in (40, 100, 40):
out = m.predict(close.iloc[-w:], horizon=1)
self.assertEqual(len(out), 1)
self.assertTrue(np.isfinite(out[0]), f'TimesFM NaN at w={w}: {out[0]}')
class TestInputContractGuards(unittest.TestCase):
def test_nan_history_raises_for_all_models(self):
from models.registry import fresh_model
rng = np.random.default_rng(3)
good = pd.Series(100 + np.cumsum(rng.normal(0, 0.5, 200)))
for name in ('ARIMA', 'Auto-ARIMA', 'ARIMA-GARCH', 'Moirai'):
m = fresh_model(name)
h = good.copy()
h.iloc[50] = np.nan
with self.assertRaises(ValueError, msg=name):
m.predict(h, horizon=1)
m = fresh_model('TimesFM')
h = good.copy()
h.iloc[50] = np.nan
with self.assertRaises(ValueError):
m.predict(h, horizon=1)
def test_horizon_below_one_raises_for_all_models(self):
from models.registry import fresh_model
rng = np.random.default_rng(3)
good = pd.Series(100 + np.cumsum(rng.normal(0, 0.5, 200)))
for name in ('ARIMA', 'Auto-ARIMA', 'ARIMA-GARCH', 'Moirai', 'TimesFM'):
m = fresh_model(name)
for bad in (0, -1):
with self.assertRaises(ValueError, msg=f'{name} horizon={bad}'):
m.predict(good, horizon=bad)
if __name__ == '__main__':
unittest.main()