twiga 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- twiga/__init__.py +102 -0
- twiga/core/__init__.py +58 -0
- twiga/core/backtester.py +499 -0
- twiga/core/config/__init__.py +27 -0
- twiga/core/config/base.py +87 -0
- twiga/core/config/characterisation.py +120 -0
- twiga/core/config/conformal.py +124 -0
- twiga/core/config/data.py +100 -0
- twiga/core/config/forecaster.py +98 -0
- twiga/core/config/neural.py +449 -0
- twiga/core/config/search_space.py +153 -0
- twiga/core/data/__init__.py +28 -0
- twiga/core/data/autores.py +193 -0
- twiga/core/data/characterisation.py +1291 -0
- twiga/core/data/diff.py +107 -0
- twiga/core/data/feature.py +269 -0
- twiga/core/data/loader.py +136 -0
- twiga/core/data/pipeline.py +362 -0
- twiga/core/data/processing.py +1404 -0
- twiga/core/data/relevance.py +135 -0
- twiga/core/data/selection.py +199 -0
- twiga/core/data/temporal.py +154 -0
- twiga/core/exceptions.py +55 -0
- twiga/core/explain/__init__.py +14 -0
- twiga/core/explain/shap_explainer.py +383 -0
- twiga/core/metrics/__init__.py +107 -0
- twiga/core/metrics/_core.py +67 -0
- twiga/core/metrics/interval.py +453 -0
- twiga/core/metrics/parametric.py +144 -0
- twiga/core/metrics/point.py +975 -0
- twiga/core/metrics/prob.py +371 -0
- twiga/core/metrics/quantile.py +735 -0
- twiga/core/metrics/stats.py +265 -0
- twiga/core/plot/__init__.py +84 -0
- twiga/core/plot/_constants.py +57 -0
- twiga/core/plot/distribution.py +396 -0
- twiga/core/plot/exploration.py +399 -0
- twiga/core/plot/gt.py +311 -0
- twiga/core/plot/matplot_theme.py +224 -0
- twiga/core/plot/plotly.py +302 -0
- twiga/core/plot/residuals.py +457 -0
- twiga/core/plot/stats.py +530 -0
- twiga/core/plot/theme.py +170 -0
- twiga/core/plot/timeseries.py +835 -0
- twiga/core/settings.py +77 -0
- twiga/core/stats/__init__.py +64 -0
- twiga/core/stats/ami.py +147 -0
- twiga/core/stats/association.py +122 -0
- twiga/core/stats/autocorr.py +195 -0
- twiga/core/stats/entropy.py +220 -0
- twiga/core/stats/mutual_information.py +68 -0
- twiga/core/stats/ppscore.py +139 -0
- twiga/core/stats/residual.py +128 -0
- twiga/core/stats/seasonality.py +159 -0
- twiga/core/stats/selection.py +87 -0
- twiga/core/stats/stationarity.py +94 -0
- twiga/core/stats/utils.py +448 -0
- twiga/core/stats/xicorr.py +144 -0
- twiga/core/utils/__init__.py +5 -0
- twiga/core/utils/logger.py +379 -0
- twiga/distributions/__init__.py +28 -0
- twiga/distributions/conformal/__init__.py +14 -0
- twiga/distributions/conformal/base.py +198 -0
- twiga/distributions/conformal/common.py +0 -0
- twiga/distributions/conformal/core.py +25 -0
- twiga/distributions/conformal/cqr.py +90 -0
- twiga/distributions/conformal/crc.py +86 -0
- twiga/distributions/conformal/util +0 -0
- twiga/distributions/ml/__init__.py +0 -0
- twiga/distributions/ml/utils.py +117 -0
- twiga/distributions/nn/__init__.py +55 -0
- twiga/distributions/nn/custom_loss.py +334 -0
- twiga/distributions/nn/fpquantile.py +426 -0
- twiga/distributions/nn/parametric.py +701 -0
- twiga/distributions/nn/quantile.py +244 -0
- twiga/distributions/nn/residual_conformal.py +273 -0
- twiga/forecaster/__init__.py +18 -0
- twiga/forecaster/abstract.py +904 -0
- twiga/forecaster/base.py +1045 -0
- twiga/forecaster/core.py +363 -0
- twiga/forecaster/ensemble.py +76 -0
- twiga/forecaster/registry.py +99 -0
- twiga/forecaster/result.py +362 -0
- twiga/forecaster/utils.py +312 -0
- twiga/models/__init__.py +94 -0
- twiga/models/baseline/__init__.py +20 -0
- twiga/models/baseline/context_parrot_model.py +330 -0
- twiga/models/baseline/drift_model.py +179 -0
- twiga/models/baseline/naive_model.py +222 -0
- twiga/models/baseline/seasonal_naive_model.py +285 -0
- twiga/models/baseline/window_average_model.py +206 -0
- twiga/models/ml/__init__.py +80 -0
- twiga/models/ml/catboost_model.py +98 -0
- twiga/models/ml/core/__init__.py +0 -0
- twiga/models/ml/core/base_regressor.py +234 -0
- twiga/models/ml/gausscatboost_model.py +290 -0
- twiga/models/ml/lightgbm_model.py +197 -0
- twiga/models/ml/lineareg_model.py +74 -0
- twiga/models/ml/ngboostexponential_model.py +160 -0
- twiga/models/ml/ngboostlognormal_model.py +166 -0
- twiga/models/ml/ngboostnormal_model.py +152 -0
- twiga/models/ml/prob/__init__.py +0 -0
- twiga/models/ml/prob/base_ngboost.py +149 -0
- twiga/models/ml/prob/base_quantile.py +318 -0
- twiga/models/ml/qrcatboost_model.py +96 -0
- twiga/models/ml/qrlightgbm_model.py +71 -0
- twiga/models/ml/qrrandomforest_model.py +196 -0
- twiga/models/ml/qrxgboost_model.py +85 -0
- twiga/models/ml/randomforest_model.py +115 -0
- twiga/models/ml/xgboost_model.py +162 -0
- twiga/models/nn/__init__.py +140 -0
- twiga/models/nn/core/__init__.py +46 -0
- twiga/models/nn/core/base.py +376 -0
- twiga/models/nn/core/base_arch.py +191 -0
- twiga/models/nn/core/base_model.py +541 -0
- twiga/models/nn/core/defaults.py +305 -0
- twiga/models/nn/core/embedding.py +673 -0
- twiga/models/nn/core/linear.py +729 -0
- twiga/models/nn/core/scheduler.py +0 -0
- twiga/models/nn/mlpf_model.py +187 -0
- twiga/models/nn/mlpfbeta_model.py +139 -0
- twiga/models/nn/mlpfcrc_model.py +142 -0
- twiga/models/nn/mlpffpqr_model.py +207 -0
- twiga/models/nn/mlpfgamma_model.py +139 -0
- twiga/models/nn/mlpflaplace_model.py +128 -0
- twiga/models/nn/mlpflognormal_model.py +132 -0
- twiga/models/nn/mlpfnormal_model.py +127 -0
- twiga/models/nn/mlpfqr_model.py +381 -0
- twiga/models/nn/mlpfstudentt_model.py +135 -0
- twiga/models/nn/mlpgaf_model.py +224 -0
- twiga/models/nn/mlpgafbeta_model.py +142 -0
- twiga/models/nn/mlpgafcrc_model.py +150 -0
- twiga/models/nn/mlpgaffpqr_model.py +169 -0
- twiga/models/nn/mlpgafgamma_model.py +142 -0
- twiga/models/nn/mlpgaflaplace_model.py +132 -0
- twiga/models/nn/mlpgaflognormal_model.py +132 -0
- twiga/models/nn/mlpgafnormal_model.py +139 -0
- twiga/models/nn/mlpgafqr_model.py +166 -0
- twiga/models/nn/mlpgafstudentt_model.py +137 -0
- twiga/models/nn/mlpgam_model.py +169 -0
- twiga/models/nn/mlpgambeta_model.py +138 -0
- twiga/models/nn/mlpgamcrc_model.py +136 -0
- twiga/models/nn/mlpgamfpqr_model.py +163 -0
- twiga/models/nn/mlpgamgamma_model.py +138 -0
- twiga/models/nn/mlpgamlaplace_model.py +128 -0
- twiga/models/nn/mlpgamlognormal_model.py +128 -0
- twiga/models/nn/mlpgamnormal_model.py +134 -0
- twiga/models/nn/mlpgamqr_model.py +162 -0
- twiga/models/nn/mlpgamstudentt_model.py +133 -0
- twiga/models/nn/net/__init__.py +17 -0
- twiga/models/nn/net/ganf/__init__.py +0 -0
- twiga/models/nn/net/ganf/gates.py +73 -0
- twiga/models/nn/net/ganf/groups.py +170 -0
- twiga/models/nn/net/ganf/importance.py +829 -0
- twiga/models/nn/net/ganf/losses.py +159 -0
- twiga/models/nn/net/mlpf.py +414 -0
- twiga/models/nn/net/mlpfgaf.py +589 -0
- twiga/models/nn/net/mlpfgam.py +468 -0
- twiga/models/nn/net/nhits.py +723 -0
- twiga/models/nn/net/rnn.py +257 -0
- twiga/models/nn/nhits_model.py +306 -0
- twiga/models/nn/nhitsbeta_model.py +98 -0
- twiga/models/nn/nhitscrc_model.py +122 -0
- twiga/models/nn/nhitsgamma_model.py +98 -0
- twiga/models/nn/nhitslaplace_model.py +98 -0
- twiga/models/nn/nhitslognormal_model.py +98 -0
- twiga/models/nn/nhitsnormal_model.py +98 -0
- twiga/models/nn/nhitsqr_model.py +239 -0
- twiga/models/nn/nhitsstudentt_model.py +101 -0
- twiga/models/nn/prob/__init__.py +81 -0
- twiga/models/nn/prob/core.py +201 -0
- twiga/models/nn/prob/mlpf_crc.py +120 -0
- twiga/models/nn/prob/mlpf_parametric.py +456 -0
- twiga/models/nn/prob/mlpffpqr.py +144 -0
- twiga/models/nn/prob/mlpfqr.py +143 -0
- twiga/models/nn/prob/mlpgaf_crc.py +130 -0
- twiga/models/nn/prob/mlpgaf_parametric.py +499 -0
- twiga/models/nn/prob/mlpgaffpqr.py +162 -0
- twiga/models/nn/prob/mlpgafqr.py +162 -0
- twiga/models/nn/prob/mlpgam_crc.py +118 -0
- twiga/models/nn/prob/mlpgam_parametric.py +453 -0
- twiga/models/nn/prob/mlpgamfpqr.py +143 -0
- twiga/models/nn/prob/mlpgamqr.py +142 -0
- twiga/models/nn/prob/nhits_crc.py +107 -0
- twiga/models/nn/prob/nhits_parametric.py +412 -0
- twiga/models/nn/prob/nhitsfpqr.py +132 -0
- twiga/models/nn/prob/nhitsqr.py +131 -0
- twiga/models/nn/prob/rnn_parametric.py +500 -0
- twiga/models/nn/prob/rnn_qr.py +242 -0
- twiga/models/nn/rnn_model.py +159 -0
- twiga/models/nn/rnnbeta_model.py +108 -0
- twiga/models/nn/rnngamma_model.py +108 -0
- twiga/models/nn/rnnlaplace_model.py +108 -0
- twiga/models/nn/rnnlognormal_model.py +108 -0
- twiga/models/nn/rnnnormal_model.py +108 -0
- twiga/models/nn/rnnqr_model.py +251 -0
- twiga/models/nn/rnnstudentt_model.py +108 -0
- twiga/pipeline/__init__.py +52 -0
- twiga/pipeline/flows.py +441 -0
- twiga/pipeline/schemas.py +98 -0
- twiga/pipeline/tasks.py +350 -0
- twiga/serve/__init__.py +41 -0
- twiga/serve/app.py +408 -0
- twiga/serve/loader.py +174 -0
- twiga/serve/monitor.py +462 -0
- twiga/serve/schemas.py +263 -0
- twiga/tracking/__init__.py +22 -0
- twiga/tracking/tracker.py +314 -0
- twiga/utils/__init__.py +7 -0
- twiga/utils/data_loading.py +281 -0
- twiga/utils/run_forecast.py +757 -0
- twiga/utils/utils.py +36 -0
- twiga-0.1.0.dist-info/METADATA +467 -0
- twiga-0.1.0.dist-info/RECORD +216 -0
- twiga-0.1.0.dist-info/WHEEL +4 -0
- twiga-0.1.0.dist-info/licenses/LICENSE +201 -0
twiga/__init__.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
"""Twiga - point and probabilistic time series forecasting.
|
|
2
|
+
|
|
3
|
+
Public API
|
|
4
|
+
----------
|
|
5
|
+
The symbols exported here are the **stable public API**. Everything else
|
|
6
|
+
(submodules, internal helpers, model classes) may change without notice.
|
|
7
|
+
|
|
8
|
+
Stable
|
|
9
|
+
~~~~~~
|
|
10
|
+
- :class:`TwigaForecaster` - main entry point
|
|
11
|
+
- :class:`DataPipelineConfig` - data pipeline configuration
|
|
12
|
+
- :class:`ForecasterConfig` - training configuration
|
|
13
|
+
- :class:`BaseModelConfig` - base model configuration
|
|
14
|
+
- :class:`ConformalConfig` - conformal prediction configuration
|
|
15
|
+
- :func:`get_model` - registry lookup by name
|
|
16
|
+
- :func:`evaluate_forecast` - evaluation helper
|
|
17
|
+
|
|
18
|
+
Experimental (may change in minor versions)
|
|
19
|
+
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
|
20
|
+
- :class:`BaseMLPConfig`, :class:`NeuralModelConfig`, :class:`BaseSearchSpace`
|
|
21
|
+
- :class:`ForecastResult`, :class:`ForecastCollection`, :class:`ForecastKind`
|
|
22
|
+
|
|
23
|
+
Optional submodule imports
|
|
24
|
+
~~~~~~~~~~~~~~~~~~~~~~~~~~
|
|
25
|
+
Tree-boosting models require ``pip install 'twiga[gbtree]'``::
|
|
26
|
+
|
|
27
|
+
from twiga.models.ml import LIGHTGBMConfig, LIGHTGBMModel
|
|
28
|
+
|
|
29
|
+
Neural network models require ``pip install 'twiga[nn]'``::
|
|
30
|
+
|
|
31
|
+
from twiga.models.nn import MLPFConfig, MLPFModel
|
|
32
|
+
|
|
33
|
+
SHAP explainability requires ``pip install 'twiga[explain]'``::
|
|
34
|
+
|
|
35
|
+
from twiga.core.explain import ShapExplainer
|
|
36
|
+
|
|
37
|
+
Plotting requires ``pip install 'twiga[plots]'``::
|
|
38
|
+
|
|
39
|
+
from twiga.core.plot import plot_forecast
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
from importlib.metadata import PackageNotFoundError, version
|
|
43
|
+
|
|
44
|
+
from twiga.core.config import (
|
|
45
|
+
BaseMLPConfig,
|
|
46
|
+
BaseModelConfig,
|
|
47
|
+
BaseSearchSpace,
|
|
48
|
+
ConformalConfig,
|
|
49
|
+
DataPipelineConfig,
|
|
50
|
+
ForecasterConfig,
|
|
51
|
+
NeuralModelConfig,
|
|
52
|
+
)
|
|
53
|
+
from twiga.core.exceptions import (
|
|
54
|
+
ConfigurationError,
|
|
55
|
+
MissingExtraError,
|
|
56
|
+
NotFittedError,
|
|
57
|
+
PipelineError,
|
|
58
|
+
TwigaError,
|
|
59
|
+
require_extra,
|
|
60
|
+
)
|
|
61
|
+
from twiga.core.metrics import evaluate_forecast
|
|
62
|
+
from twiga.core.utils import configure, get_logger
|
|
63
|
+
from twiga.forecaster.core import TwigaForecaster
|
|
64
|
+
from twiga.forecaster.registry import get_model
|
|
65
|
+
from twiga.forecaster.result import ForecastCollection, ForecastKind, ForecastResult
|
|
66
|
+
|
|
67
|
+
__all__ = [
|
|
68
|
+
# ── Stable public API ──────────────────────────────────────────────────
|
|
69
|
+
# Main entry point
|
|
70
|
+
"TwigaForecaster",
|
|
71
|
+
# Configuration
|
|
72
|
+
"DataPipelineConfig",
|
|
73
|
+
"ForecasterConfig",
|
|
74
|
+
"BaseModelConfig",
|
|
75
|
+
"ConformalConfig",
|
|
76
|
+
# Registry
|
|
77
|
+
"get_model",
|
|
78
|
+
# Evaluation
|
|
79
|
+
"evaluate_forecast",
|
|
80
|
+
# Exceptions
|
|
81
|
+
"TwigaError",
|
|
82
|
+
"ConfigurationError",
|
|
83
|
+
"MissingExtraError",
|
|
84
|
+
"NotFittedError",
|
|
85
|
+
"PipelineError",
|
|
86
|
+
"require_extra",
|
|
87
|
+
# Logging
|
|
88
|
+
"configure",
|
|
89
|
+
"get_logger",
|
|
90
|
+
# ── Experimental ──────────────────────────────────────────────────────
|
|
91
|
+
"BaseMLPConfig",
|
|
92
|
+
"BaseSearchSpace",
|
|
93
|
+
"NeuralModelConfig",
|
|
94
|
+
"ForecastCollection",
|
|
95
|
+
"ForecastKind",
|
|
96
|
+
"ForecastResult",
|
|
97
|
+
]
|
|
98
|
+
|
|
99
|
+
try:
|
|
100
|
+
__version__ = version("twiga")
|
|
101
|
+
except PackageNotFoundError:
|
|
102
|
+
__version__ = "0.0.1"
|
twiga/core/__init__.py
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
"""Core infrastructure: config, data pipeline, metrics, logging."""
|
|
2
|
+
|
|
3
|
+
from twiga.core.config import (
|
|
4
|
+
BaseMLPConfig,
|
|
5
|
+
BaseModelConfig,
|
|
6
|
+
BaseSearchSpace,
|
|
7
|
+
ConformalConfig,
|
|
8
|
+
DataPipelineConfig,
|
|
9
|
+
ForecasterConfig,
|
|
10
|
+
NeuralModelConfig,
|
|
11
|
+
)
|
|
12
|
+
from twiga.core.exceptions import (
|
|
13
|
+
ConfigurationError,
|
|
14
|
+
MissingExtraError,
|
|
15
|
+
NotFittedError,
|
|
16
|
+
PipelineError,
|
|
17
|
+
TwigaError,
|
|
18
|
+
require_extra,
|
|
19
|
+
)
|
|
20
|
+
from twiga.core.metrics import (
|
|
21
|
+
evaluate_forecast,
|
|
22
|
+
evaluate_interval_forecast,
|
|
23
|
+
evaluate_parametric_forecast,
|
|
24
|
+
evaluate_point_forecast,
|
|
25
|
+
evaluate_quantile_forecast,
|
|
26
|
+
evaluate_samples_forecast,
|
|
27
|
+
)
|
|
28
|
+
from twiga.core.utils import LogContext, configure, get_logger, log_section
|
|
29
|
+
|
|
30
|
+
__all__ = [
|
|
31
|
+
# Exceptions
|
|
32
|
+
"TwigaError",
|
|
33
|
+
"ConfigurationError",
|
|
34
|
+
"MissingExtraError",
|
|
35
|
+
"NotFittedError",
|
|
36
|
+
"PipelineError",
|
|
37
|
+
"require_extra",
|
|
38
|
+
# Config
|
|
39
|
+
"BaseMLPConfig",
|
|
40
|
+
"BaseModelConfig",
|
|
41
|
+
"BaseSearchSpace",
|
|
42
|
+
"ConformalConfig",
|
|
43
|
+
"DataPipelineConfig",
|
|
44
|
+
"ForecasterConfig",
|
|
45
|
+
"NeuralModelConfig",
|
|
46
|
+
# Metrics
|
|
47
|
+
"evaluate_forecast",
|
|
48
|
+
"evaluate_interval_forecast",
|
|
49
|
+
"evaluate_parametric_forecast",
|
|
50
|
+
"evaluate_point_forecast",
|
|
51
|
+
"evaluate_quantile_forecast",
|
|
52
|
+
"evaluate_samples_forecast",
|
|
53
|
+
# Logging
|
|
54
|
+
"LogContext",
|
|
55
|
+
"configure",
|
|
56
|
+
"get_logger",
|
|
57
|
+
"log_section",
|
|
58
|
+
]
|
twiga/core/backtester.py
ADDED
|
@@ -0,0 +1,499 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
|
|
3
|
+
from dateutil.relativedelta import relativedelta # type: ignore[import-untyped]
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
from twiga.core.utils.logger import get_logger
|
|
8
|
+
|
|
9
|
+
log = get_logger(__name__)
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class SplitState:
|
|
13
|
+
"""Holds time periods for training and forecasting splits.
|
|
14
|
+
|
|
15
|
+
Attributes:
|
|
16
|
+
train_start (pd.Timestamp): Start timestamp of training period.
|
|
17
|
+
train_end (pd.Timestamp): End timestamp of training period.
|
|
18
|
+
forecast_start (pd.Timestamp): Start timestamp of forecast period.
|
|
19
|
+
forecast_end (pd.Timestamp): End timestamp of forecast period.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
def __init__(self, train_start, train_end, forecast_start, forecast_end):
|
|
23
|
+
self.train_start = train_start
|
|
24
|
+
self.train_end = train_end
|
|
25
|
+
self.forecast_start = forecast_start
|
|
26
|
+
self.forecast_end = forecast_end
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class TimeBasedSplit(ABC):
|
|
30
|
+
"""Abstract base class implementing core time-based splitting logic.
|
|
31
|
+
|
|
32
|
+
This class validates split parameters and provides properties to compute
|
|
33
|
+
time deltas for the training period, forecast period, gap, and stride.
|
|
34
|
+
|
|
35
|
+
Attributes:
|
|
36
|
+
split_freq (str): Time unit for splits (e.g., 'days', 'months').
|
|
37
|
+
train_size (int): Training period length in `split_freq` units.
|
|
38
|
+
test_size (int): Forecast period length in `split_freq` units.
|
|
39
|
+
gap (int): Gap between train and forecast periods.
|
|
40
|
+
stride (int): Step size between splits.
|
|
41
|
+
window (str): Window type ('rolling' or 'expanding').
|
|
42
|
+
"""
|
|
43
|
+
|
|
44
|
+
def __init__(
|
|
45
|
+
self,
|
|
46
|
+
split_freq: str,
|
|
47
|
+
train_size: int,
|
|
48
|
+
test_size: int,
|
|
49
|
+
gap: int = 0,
|
|
50
|
+
stride: int | None = None,
|
|
51
|
+
window: str = "rolling",
|
|
52
|
+
) -> None:
|
|
53
|
+
"""Initialize time-based split parameters.
|
|
54
|
+
|
|
55
|
+
Args:
|
|
56
|
+
split_freq (str): Time unit for splits (e.g., 'days', 'months').
|
|
57
|
+
train_size (int): Training period length in split_freq units.
|
|
58
|
+
test_size (int): Forecast period length in split_freq units.
|
|
59
|
+
gap (int): Gap between training and forecast periods (default: 0).
|
|
60
|
+
stride (int, optional): Step size between splits (default: test_size).
|
|
61
|
+
window (str): Window type ('rolling' or 'expanding') (default: 'rolling').
|
|
62
|
+
|
|
63
|
+
Raises:
|
|
64
|
+
ValueError: If any parameter is invalid. In particular, train_size must be a
|
|
65
|
+
positive integer that is greater than or equal to test_size.
|
|
66
|
+
"""
|
|
67
|
+
self.split_freq = split_freq
|
|
68
|
+
self.train_size = train_size
|
|
69
|
+
self.test_size = test_size
|
|
70
|
+
self.gap = gap
|
|
71
|
+
self.stride = stride if stride is not None else test_size
|
|
72
|
+
self.window = window
|
|
73
|
+
self._validate_arguments()
|
|
74
|
+
|
|
75
|
+
def _validate_arguments(self) -> None:
|
|
76
|
+
"""Validate input parameters meet requirements."""
|
|
77
|
+
valid_frequencies = ["days", "minutes", "hours", "weeks", "months", "years"]
|
|
78
|
+
if self.split_freq not in valid_frequencies:
|
|
79
|
+
raise ValueError(f"split_freq must be one of {valid_frequencies}. Got '{self.split_freq}'.")
|
|
80
|
+
if self.window not in ["rolling", "expanding"]:
|
|
81
|
+
raise ValueError("window must be either 'rolling' or 'expanding'.")
|
|
82
|
+
if not isinstance(self.gap, int) or self.gap < 0:
|
|
83
|
+
raise ValueError("gap must be a non-negative integer.")
|
|
84
|
+
if not isinstance(self.train_size, int) or self.train_size <= 0:
|
|
85
|
+
raise ValueError("train_size must be a positive integer.")
|
|
86
|
+
if not isinstance(self.test_size, int) or self.test_size <= 0:
|
|
87
|
+
raise ValueError("test_size must be a positive integer.")
|
|
88
|
+
if self.train_size < self.test_size:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
f"train_size must be greater than or equal to test_size. "
|
|
91
|
+
f"Got train_size = {self.train_size} and test_size = {self.test_size}.",
|
|
92
|
+
)
|
|
93
|
+
if not isinstance(self.stride, int) or self.stride <= 0:
|
|
94
|
+
raise ValueError("stride must be a positive integer.")
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def train_delta(self) -> relativedelta:
|
|
98
|
+
"""Calculate training period duration."""
|
|
99
|
+
return relativedelta(**{self.split_freq: self.train_size})
|
|
100
|
+
|
|
101
|
+
@property
|
|
102
|
+
def forecast_delta(self) -> relativedelta:
|
|
103
|
+
"""Calculate forecast period duration."""
|
|
104
|
+
return relativedelta(**{self.split_freq: self.test_size})
|
|
105
|
+
|
|
106
|
+
@property
|
|
107
|
+
def gap_delta(self) -> relativedelta:
|
|
108
|
+
"""Calculate gap duration."""
|
|
109
|
+
return relativedelta(**{self.split_freq: self.gap})
|
|
110
|
+
|
|
111
|
+
@property
|
|
112
|
+
def stride_delta(self) -> relativedelta:
|
|
113
|
+
"""Calculate stride duration."""
|
|
114
|
+
return relativedelta(**{self.split_freq: self.stride})
|
|
115
|
+
|
|
116
|
+
def _splits_from_period(self, time_start, time_end):
|
|
117
|
+
"""Generate splits between start and end times.
|
|
118
|
+
|
|
119
|
+
Args:
|
|
120
|
+
time_start (pd.Timestamp): Start timestamp.
|
|
121
|
+
time_end (pd.Timestamp): End timestamp.
|
|
122
|
+
|
|
123
|
+
Yields:
|
|
124
|
+
SplitState: A split state containing train and forecast period boundaries.
|
|
125
|
+
|
|
126
|
+
Raises:
|
|
127
|
+
ValueError: If time_start is not before time_end.
|
|
128
|
+
"""
|
|
129
|
+
if time_start >= time_end:
|
|
130
|
+
raise ValueError("time_start must be before time_end.")
|
|
131
|
+
|
|
132
|
+
is_rolling = self.window == "rolling"
|
|
133
|
+
train_start = time_start
|
|
134
|
+
train_end = train_start + self.train_delta
|
|
135
|
+
forecast_start = train_end + self.gap_delta
|
|
136
|
+
forecast_end = forecast_start + self.forecast_delta
|
|
137
|
+
|
|
138
|
+
while forecast_end <= time_end:
|
|
139
|
+
yield SplitState(train_start, train_end, forecast_start, forecast_end)
|
|
140
|
+
if is_rolling:
|
|
141
|
+
train_start += self.stride_delta
|
|
142
|
+
train_end += self.stride_delta
|
|
143
|
+
forecast_start += self.stride_delta
|
|
144
|
+
forecast_end += self.stride_delta
|
|
145
|
+
|
|
146
|
+
@abstractmethod
|
|
147
|
+
def split(self, data: pd.DataFrame):
|
|
148
|
+
"""Generate train/test splits from data."""
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
class TimeBasedCV(TimeBasedSplit):
|
|
152
|
+
"""Concrete time-based cross-validation implementation for pandas DataFrames.
|
|
153
|
+
|
|
154
|
+
This class creates splits based on a datetime column and returns train/test indices,
|
|
155
|
+
along with the corresponding time periods.
|
|
156
|
+
|
|
157
|
+
Attributes:
|
|
158
|
+
date_column (str): Name of the timestamp column in the DataFrame.
|
|
159
|
+
num_splits (int): Optional number of splits to generate.
|
|
160
|
+
"""
|
|
161
|
+
|
|
162
|
+
def __init__(
|
|
163
|
+
self,
|
|
164
|
+
split_freq,
|
|
165
|
+
test_size,
|
|
166
|
+
train_size=None,
|
|
167
|
+
gap=0,
|
|
168
|
+
stride=None,
|
|
169
|
+
window="rolling",
|
|
170
|
+
date_column="timestamp",
|
|
171
|
+
num_splits=None,
|
|
172
|
+
):
|
|
173
|
+
if (train_size is None) and (num_splits is None):
|
|
174
|
+
raise ValueError("Either train_size or num_splits must be provided, not both None.")
|
|
175
|
+
self.split_freq = split_freq
|
|
176
|
+
self.test_size = test_size
|
|
177
|
+
self.train_size = train_size if train_size is not None else 0 # Set temporarily
|
|
178
|
+
self.gap = gap
|
|
179
|
+
self.stride = stride if stride is not None else test_size
|
|
180
|
+
self.window = window
|
|
181
|
+
self.date_column = date_column
|
|
182
|
+
self.num_splits = num_splits
|
|
183
|
+
self._split_scheme = {}
|
|
184
|
+
self.time_values = None
|
|
185
|
+
if train_size is not None:
|
|
186
|
+
super().__init__(
|
|
187
|
+
split_freq=split_freq,
|
|
188
|
+
train_size=train_size,
|
|
189
|
+
test_size=test_size,
|
|
190
|
+
gap=gap,
|
|
191
|
+
stride=self.stride,
|
|
192
|
+
window=window,
|
|
193
|
+
)
|
|
194
|
+
|
|
195
|
+
def duration_in_units(self, start: pd.Timestamp, end: pd.Timestamp, split_freq: str) -> int:
|
|
196
|
+
"""Compute the duration between start and end in the specified split_freq units.
|
|
197
|
+
|
|
198
|
+
For 'days', 'minutes', 'hours', and 'weeks', a simple conversion based on timedelta is used.
|
|
199
|
+
For 'months' and 'years', relativedelta is used to account for variable lengths.
|
|
200
|
+
|
|
201
|
+
Args:
|
|
202
|
+
start (pd.Timestamp): Start timestamp.
|
|
203
|
+
end (pd.Timestamp): End timestamp.
|
|
204
|
+
split_freq (str): One of 'days', 'minutes', 'hours', 'weeks', 'months', or 'years'.
|
|
205
|
+
|
|
206
|
+
Returns:
|
|
207
|
+
int: Duration in the specified units.
|
|
208
|
+
|
|
209
|
+
Raises:
|
|
210
|
+
ValueError: If split_freq is unsupported.
|
|
211
|
+
"""
|
|
212
|
+
diff = end - start
|
|
213
|
+
conversion = {
|
|
214
|
+
"days": diff.days,
|
|
215
|
+
"minutes": int(diff.total_seconds() / 60),
|
|
216
|
+
"hours": int(diff.total_seconds() / 3600),
|
|
217
|
+
"weeks": diff.days // 7,
|
|
218
|
+
}
|
|
219
|
+
if split_freq in conversion:
|
|
220
|
+
return conversion[split_freq]
|
|
221
|
+
if split_freq == "months":
|
|
222
|
+
rd = relativedelta(end, start)
|
|
223
|
+
return rd.years * 12 + rd.months
|
|
224
|
+
if split_freq == "years":
|
|
225
|
+
return relativedelta(end, start).years
|
|
226
|
+
raise ValueError(f"Unsupported split_freq: {split_freq}")
|
|
227
|
+
|
|
228
|
+
def _validate_train_size(self) -> None:
|
|
229
|
+
"""Ensure that the computed train_size is valid and at least as long as test_size."""
|
|
230
|
+
if not isinstance(self.train_size, int) or self.train_size <= 0:
|
|
231
|
+
raise ValueError(f"Calculated train_size must be a positive integer. Got {self.train_size}.")
|
|
232
|
+
if self.train_size < self.test_size:
|
|
233
|
+
raise ValueError(
|
|
234
|
+
f"Calculated train_size must be greater than or equal to test_size. "
|
|
235
|
+
f"Got train_size = {self.train_size}, test_size = {self.test_size}.",
|
|
236
|
+
)
|
|
237
|
+
|
|
238
|
+
def set_split_scheme(
|
|
239
|
+
self,
|
|
240
|
+
time_values: pd.Series,
|
|
241
|
+
start_dt: pd.Timestamp | None = None,
|
|
242
|
+
end_dt: pd.Timestamp | None = None,
|
|
243
|
+
) -> None:
|
|
244
|
+
"""Calculate split indices from time series data.
|
|
245
|
+
|
|
246
|
+
The method sorts the datetime series, determines the time range to use,
|
|
247
|
+
and computes indices for training and forecast periods based on the provided
|
|
248
|
+
parameters. If num_splits is set, it adjusts train_size accordingly.
|
|
249
|
+
|
|
250
|
+
Args:
|
|
251
|
+
time_values (pd.Series): Datetime series to split.
|
|
252
|
+
start_dt (pd.Timestamp, optional): Override start time (default: min of series).
|
|
253
|
+
end_dt (pd.Timestamp, optional): Override end time (default: max of series).
|
|
254
|
+
|
|
255
|
+
Raises:
|
|
256
|
+
ValueError: If time_values is not a datetime series.
|
|
257
|
+
"""
|
|
258
|
+
if not pd.api.types.is_datetime64_any_dtype(time_values):
|
|
259
|
+
raise ValueError("time_values must be a datetime series.")
|
|
260
|
+
|
|
261
|
+
# Sort data by time
|
|
262
|
+
indices = np.arange(len(time_values))
|
|
263
|
+
sorted_order = np.argsort(time_values)
|
|
264
|
+
self.time_values = time_values.iloc[sorted_order]
|
|
265
|
+
indices = indices[sorted_order]
|
|
266
|
+
|
|
267
|
+
# Determine time range
|
|
268
|
+
time_start = start_dt or self.time_values.min()
|
|
269
|
+
time_end = end_dt or self.time_values.max()
|
|
270
|
+
|
|
271
|
+
# Adjust parameters if num_splits is specified.
|
|
272
|
+
if self.num_splits is not None:
|
|
273
|
+
total_duration = self.duration_in_units(time_start, time_end, self.split_freq)
|
|
274
|
+
max_possible_splits = (total_duration - 2 * self.test_size - self.gap) / self.stride
|
|
275
|
+
if self.num_splits >= max_possible_splits:
|
|
276
|
+
log.warning(
|
|
277
|
+
"Not enough time for %d splits. Reducing num_splits to maximum possible: %d.",
|
|
278
|
+
self.num_splits,
|
|
279
|
+
max_possible_splits,
|
|
280
|
+
)
|
|
281
|
+
self.num_splits = int(max_possible_splits)
|
|
282
|
+
self.train_size = total_duration - self.test_size - self.gap - (self.num_splits - 1) * self.stride
|
|
283
|
+
self._validate_train_size()
|
|
284
|
+
|
|
285
|
+
self._split_scheme = {}
|
|
286
|
+
for i, split in enumerate(self._splits_from_period(time_start, time_end)):
|
|
287
|
+
train_mask = (self.time_values >= split.train_start) & (self.time_values < split.train_end)
|
|
288
|
+
test_mask = (self.time_values >= split.forecast_start) & (self.time_values < split.forecast_end)
|
|
289
|
+
|
|
290
|
+
if train_mask.any() and test_mask.any():
|
|
291
|
+
self._split_scheme[i] = {
|
|
292
|
+
"train_idx": indices[train_mask],
|
|
293
|
+
"test_idx": indices[test_mask],
|
|
294
|
+
"train_period": (split.train_start, split.train_end),
|
|
295
|
+
"test_period": (split.forecast_start, split.forecast_end),
|
|
296
|
+
}
|
|
297
|
+
|
|
298
|
+
if self.num_splits is not None and len(self._split_scheme) != self.num_splits:
|
|
299
|
+
log.warning(
|
|
300
|
+
"Requested %d splits but generated %d. Time range: %s to %s",
|
|
301
|
+
self.num_splits,
|
|
302
|
+
len(self._split_scheme),
|
|
303
|
+
time_start,
|
|
304
|
+
time_end,
|
|
305
|
+
)
|
|
306
|
+
self.num_splits = len(self._split_scheme)
|
|
307
|
+
|
|
308
|
+
def get_scheme(self) -> dict:
|
|
309
|
+
"""Return the current split configuration.
|
|
310
|
+
|
|
311
|
+
Returns:
|
|
312
|
+
dict: A dictionary containing the train/test split indices and periods.
|
|
313
|
+
|
|
314
|
+
Raises:
|
|
315
|
+
ValueError: If the split scheme has not been initialized.
|
|
316
|
+
"""
|
|
317
|
+
if not self._split_scheme:
|
|
318
|
+
raise ValueError("Split scheme not initialized. Call set_split_scheme() first.")
|
|
319
|
+
return self._split_scheme.copy()
|
|
320
|
+
|
|
321
|
+
def split(self, data: pd.DataFrame, start_dt: pd.Timestamp | None = None, end_dt: pd.Timestamp | None = None):
|
|
322
|
+
"""Generate validated train/test splits.
|
|
323
|
+
|
|
324
|
+
Args:
|
|
325
|
+
data (pd.DataFrame): DataFrame containing the date column.
|
|
326
|
+
start_dt (pd.Timestamp, optional): Override split start time.
|
|
327
|
+
end_dt (pd.Timestamp, optional): Override split end time.
|
|
328
|
+
|
|
329
|
+
Yields:
|
|
330
|
+
tuple: (train_df, test_df, scheme, split_key).
|
|
331
|
+
|
|
332
|
+
Raises:
|
|
333
|
+
ValueError: If the required date column is missing or if the computed indices exceed data bounds.
|
|
334
|
+
"""
|
|
335
|
+
if self.date_column not in data.columns:
|
|
336
|
+
raise ValueError(f"Missing time column: '{self.date_column}' in the data.")
|
|
337
|
+
|
|
338
|
+
if not self._split_scheme:
|
|
339
|
+
self.set_split_scheme(data[self.date_column], start_dt, end_dt)
|
|
340
|
+
|
|
341
|
+
max_index = max(np.max(np.concatenate([s["train_idx"], s["test_idx"]])) for s in self._split_scheme.values())
|
|
342
|
+
if max_index >= len(data):
|
|
343
|
+
raise ValueError(
|
|
344
|
+
f"Data has {len(data)} rows but splits require index {max_index}. Check your time series range.",
|
|
345
|
+
)
|
|
346
|
+
|
|
347
|
+
for split_key, scheme in self._split_scheme.items():
|
|
348
|
+
train_df = data.iloc[scheme["train_idx"]].reset_index(drop=True)
|
|
349
|
+
test_df = data.iloc[scheme["test_idx"]].reset_index(drop=True)
|
|
350
|
+
yield train_df, test_df, scheme, split_key
|
|
351
|
+
|
|
352
|
+
def plot_split_scheme(
|
|
353
|
+
self,
|
|
354
|
+
data: pd.DataFrame | None = None,
|
|
355
|
+
train_ratio: float = 1.0,
|
|
356
|
+
start_dt: pd.Timestamp | None = None,
|
|
357
|
+
end_dt: pd.Timestamp | None = None,
|
|
358
|
+
title: str = "Cross-validation split scheme",
|
|
359
|
+
colors: dict[str, str] | None = None,
|
|
360
|
+
alpha: float = 0.88,
|
|
361
|
+
x_ticks: int = 6,
|
|
362
|
+
font_size: int = 10,
|
|
363
|
+
line_width: float = 0.8,
|
|
364
|
+
x_axis_angle: int = 30,
|
|
365
|
+
legend_pos: str = "top",
|
|
366
|
+
):
|
|
367
|
+
"""Visualize the time series cross-validation split scheme.
|
|
368
|
+
|
|
369
|
+
Renders a Gantt-style plot with one horizontal bar per fold, colour-coded
|
|
370
|
+
by segment (Train / Val / Test), styled with the Twiga theme.
|
|
371
|
+
|
|
372
|
+
Args:
|
|
373
|
+
data: Input DataFrame containing temporal data. Used to derive the
|
|
374
|
+
split scheme when it has not been pre-computed.
|
|
375
|
+
train_ratio: Proportion of training indices used for training; the
|
|
376
|
+
remainder becomes a validation segment.
|
|
377
|
+
start_dt: Optional start timestamp passed to ``set_split_scheme``.
|
|
378
|
+
end_dt: Optional end timestamp passed to ``set_split_scheme``.
|
|
379
|
+
title: Plot title.
|
|
380
|
+
colors: Custom colour mapping for segments. Keys must be title-case:
|
|
381
|
+
``"Train"``, ``"Val"``, ``"Test"``.
|
|
382
|
+
alpha: Bar transparency (0–1).
|
|
383
|
+
x_ticks: Number of date ticks on the x-axis.
|
|
384
|
+
font_size: Base font size in points.
|
|
385
|
+
line_width: Axis line stroke width.
|
|
386
|
+
x_axis_angle: Rotation angle for x-axis tick labels.
|
|
387
|
+
legend_pos: Legend position - ``"top"``, ``"bottom"``, ``"left"``,
|
|
388
|
+
``"right"``, or ``"none"``.
|
|
389
|
+
|
|
390
|
+
Returns:
|
|
391
|
+
A Lets-Plot ``ggplot`` object.
|
|
392
|
+
|
|
393
|
+
Raises:
|
|
394
|
+
ValueError: If ``train_ratio`` is outside [0, 1] or the split scheme
|
|
395
|
+
is not initialised and no ``data`` is provided.
|
|
396
|
+
|
|
397
|
+
Example:
|
|
398
|
+
>>> splitter = TimeBasedCV(split_freq="days", test_size=5, train_size=20, date_column="date")
|
|
399
|
+
>>> splitter.set_split_scheme(data["date"])
|
|
400
|
+
>>> splitter.plot_split_scheme(data, train_ratio=0.8, title="CV Scheme")
|
|
401
|
+
"""
|
|
402
|
+
from lets_plot import (
|
|
403
|
+
aes,
|
|
404
|
+
geom_rect,
|
|
405
|
+
ggplot,
|
|
406
|
+
labs,
|
|
407
|
+
scale_fill_manual,
|
|
408
|
+
scale_x_continuous,
|
|
409
|
+
scale_y_continuous,
|
|
410
|
+
)
|
|
411
|
+
|
|
412
|
+
from twiga.core.plot import TWIGA_PALETTE, twiga_theme
|
|
413
|
+
|
|
414
|
+
if not 0 <= train_ratio <= 1:
|
|
415
|
+
raise ValueError("train_ratio must be between 0 and 1.")
|
|
416
|
+
|
|
417
|
+
if not hasattr(self, "_split_scheme") or not self._split_scheme:
|
|
418
|
+
if data is not None:
|
|
419
|
+
self.set_split_scheme(data[self.date_column], start_dt, end_dt)
|
|
420
|
+
else:
|
|
421
|
+
raise ValueError("Split scheme not initialized and no data provided.")
|
|
422
|
+
|
|
423
|
+
default_colors = {
|
|
424
|
+
"Train": TWIGA_PALETTE[0], # teal
|
|
425
|
+
"Val": TWIGA_PALETTE[5], # emerald
|
|
426
|
+
"Test": TWIGA_PALETTE[1], # amber
|
|
427
|
+
}
|
|
428
|
+
# Accept legacy lowercase keys for backward compatibility
|
|
429
|
+
color_mapping = {k.title(): v for k, v in colors.items()} if colors else default_colors
|
|
430
|
+
|
|
431
|
+
_bar_h = 0.55
|
|
432
|
+
segments = []
|
|
433
|
+
for fold_idx, scheme in self._split_scheme.items():
|
|
434
|
+
fold = fold_idx + 1
|
|
435
|
+
train_indices = scheme["train_idx"]
|
|
436
|
+
train_size = int(train_ratio * len(train_indices))
|
|
437
|
+
val_size = len(train_indices) - train_size
|
|
438
|
+
|
|
439
|
+
segments.append(
|
|
440
|
+
{
|
|
441
|
+
"fold": fold,
|
|
442
|
+
"segment": "Train",
|
|
443
|
+
"start": train_indices[0],
|
|
444
|
+
"end": train_indices[0] + train_size,
|
|
445
|
+
"ymin": fold - _bar_h / 2,
|
|
446
|
+
"ymax": fold + _bar_h / 2,
|
|
447
|
+
},
|
|
448
|
+
)
|
|
449
|
+
if val_size > 0:
|
|
450
|
+
segments.append(
|
|
451
|
+
{
|
|
452
|
+
"fold": fold,
|
|
453
|
+
"segment": "Val",
|
|
454
|
+
"start": train_indices[0] + train_size,
|
|
455
|
+
"end": train_indices[0] + train_size + val_size,
|
|
456
|
+
"ymin": fold - _bar_h / 2,
|
|
457
|
+
"ymax": fold + _bar_h / 2,
|
|
458
|
+
},
|
|
459
|
+
)
|
|
460
|
+
segments.append(
|
|
461
|
+
{
|
|
462
|
+
"fold": fold,
|
|
463
|
+
"segment": "Test",
|
|
464
|
+
"start": scheme["test_idx"][0],
|
|
465
|
+
"end": scheme["test_idx"][-1],
|
|
466
|
+
"ymin": fold - _bar_h / 2,
|
|
467
|
+
"ymax": fold + _bar_h / 2,
|
|
468
|
+
},
|
|
469
|
+
)
|
|
470
|
+
|
|
471
|
+
df = pd.DataFrame(segments)
|
|
472
|
+
|
|
473
|
+
time_values = pd.Series(self.time_values)
|
|
474
|
+
x_breaks = np.linspace(0, len(time_values) - 1, x_ticks, dtype=int).tolist()
|
|
475
|
+
x_labels = time_values.iloc[x_breaks].dt.strftime("%Y-%m-%d").tolist()
|
|
476
|
+
|
|
477
|
+
n_folds = self.num_splits
|
|
478
|
+
fold_breaks = list(range(1, n_folds + 1))
|
|
479
|
+
fold_labels = [f"Fold {i}" for i in fold_breaks]
|
|
480
|
+
|
|
481
|
+
return (
|
|
482
|
+
ggplot(df)
|
|
483
|
+
+ geom_rect(
|
|
484
|
+
aes(xmin="start", xmax="end", ymin="ymin", ymax="ymax", fill="segment"),
|
|
485
|
+
size=0,
|
|
486
|
+
alpha=alpha,
|
|
487
|
+
)
|
|
488
|
+
+ scale_fill_manual(values=color_mapping)
|
|
489
|
+
+ scale_x_continuous(breaks=x_breaks, labels=x_labels)
|
|
490
|
+
+ scale_y_continuous(breaks=fold_breaks, labels=fold_labels)
|
|
491
|
+
+ labs(title=title, x="Date", y="", fill="")
|
|
492
|
+
+ twiga_theme(
|
|
493
|
+
font_size=font_size,
|
|
494
|
+
line_width=line_width,
|
|
495
|
+
x_axis_angle=x_axis_angle,
|
|
496
|
+
legend_pos=legend_pos, # type: ignore[arg-type]
|
|
497
|
+
grid=False,
|
|
498
|
+
)
|
|
499
|
+
)
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
"""Domain-specific configuration classes.
|
|
2
|
+
|
|
3
|
+
These configs describe the problem domain (data structure, CV strategy,
|
|
4
|
+
uncertainty quantification) rather than a specific model architecture.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from pydantic import Field
|
|
8
|
+
|
|
9
|
+
from .base import BaseModelConfig
|
|
10
|
+
from .characterisation import CharacterisationConfig
|
|
11
|
+
from .conformal import ConformalConfig
|
|
12
|
+
from .data import DataPipelineConfig
|
|
13
|
+
from .forecaster import ForecasterConfig
|
|
14
|
+
from .neural import BaseMLPConfig, NeuralModelConfig
|
|
15
|
+
from .search_space import BaseSearchSpace
|
|
16
|
+
|
|
17
|
+
__all__ = [
|
|
18
|
+
"BaseMLPConfig",
|
|
19
|
+
"BaseModelConfig",
|
|
20
|
+
"BaseSearchSpace",
|
|
21
|
+
"CharacterisationConfig",
|
|
22
|
+
"ConformalConfig",
|
|
23
|
+
"DataPipelineConfig",
|
|
24
|
+
"Field",
|
|
25
|
+
"ForecasterConfig",
|
|
26
|
+
"NeuralModelConfig",
|
|
27
|
+
]
|