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.
Files changed (216) hide show
  1. twiga/__init__.py +102 -0
  2. twiga/core/__init__.py +58 -0
  3. twiga/core/backtester.py +499 -0
  4. twiga/core/config/__init__.py +27 -0
  5. twiga/core/config/base.py +87 -0
  6. twiga/core/config/characterisation.py +120 -0
  7. twiga/core/config/conformal.py +124 -0
  8. twiga/core/config/data.py +100 -0
  9. twiga/core/config/forecaster.py +98 -0
  10. twiga/core/config/neural.py +449 -0
  11. twiga/core/config/search_space.py +153 -0
  12. twiga/core/data/__init__.py +28 -0
  13. twiga/core/data/autores.py +193 -0
  14. twiga/core/data/characterisation.py +1291 -0
  15. twiga/core/data/diff.py +107 -0
  16. twiga/core/data/feature.py +269 -0
  17. twiga/core/data/loader.py +136 -0
  18. twiga/core/data/pipeline.py +362 -0
  19. twiga/core/data/processing.py +1404 -0
  20. twiga/core/data/relevance.py +135 -0
  21. twiga/core/data/selection.py +199 -0
  22. twiga/core/data/temporal.py +154 -0
  23. twiga/core/exceptions.py +55 -0
  24. twiga/core/explain/__init__.py +14 -0
  25. twiga/core/explain/shap_explainer.py +383 -0
  26. twiga/core/metrics/__init__.py +107 -0
  27. twiga/core/metrics/_core.py +67 -0
  28. twiga/core/metrics/interval.py +453 -0
  29. twiga/core/metrics/parametric.py +144 -0
  30. twiga/core/metrics/point.py +975 -0
  31. twiga/core/metrics/prob.py +371 -0
  32. twiga/core/metrics/quantile.py +735 -0
  33. twiga/core/metrics/stats.py +265 -0
  34. twiga/core/plot/__init__.py +84 -0
  35. twiga/core/plot/_constants.py +57 -0
  36. twiga/core/plot/distribution.py +396 -0
  37. twiga/core/plot/exploration.py +399 -0
  38. twiga/core/plot/gt.py +311 -0
  39. twiga/core/plot/matplot_theme.py +224 -0
  40. twiga/core/plot/plotly.py +302 -0
  41. twiga/core/plot/residuals.py +457 -0
  42. twiga/core/plot/stats.py +530 -0
  43. twiga/core/plot/theme.py +170 -0
  44. twiga/core/plot/timeseries.py +835 -0
  45. twiga/core/settings.py +77 -0
  46. twiga/core/stats/__init__.py +64 -0
  47. twiga/core/stats/ami.py +147 -0
  48. twiga/core/stats/association.py +122 -0
  49. twiga/core/stats/autocorr.py +195 -0
  50. twiga/core/stats/entropy.py +220 -0
  51. twiga/core/stats/mutual_information.py +68 -0
  52. twiga/core/stats/ppscore.py +139 -0
  53. twiga/core/stats/residual.py +128 -0
  54. twiga/core/stats/seasonality.py +159 -0
  55. twiga/core/stats/selection.py +87 -0
  56. twiga/core/stats/stationarity.py +94 -0
  57. twiga/core/stats/utils.py +448 -0
  58. twiga/core/stats/xicorr.py +144 -0
  59. twiga/core/utils/__init__.py +5 -0
  60. twiga/core/utils/logger.py +379 -0
  61. twiga/distributions/__init__.py +28 -0
  62. twiga/distributions/conformal/__init__.py +14 -0
  63. twiga/distributions/conformal/base.py +198 -0
  64. twiga/distributions/conformal/common.py +0 -0
  65. twiga/distributions/conformal/core.py +25 -0
  66. twiga/distributions/conformal/cqr.py +90 -0
  67. twiga/distributions/conformal/crc.py +86 -0
  68. twiga/distributions/conformal/util +0 -0
  69. twiga/distributions/ml/__init__.py +0 -0
  70. twiga/distributions/ml/utils.py +117 -0
  71. twiga/distributions/nn/__init__.py +55 -0
  72. twiga/distributions/nn/custom_loss.py +334 -0
  73. twiga/distributions/nn/fpquantile.py +426 -0
  74. twiga/distributions/nn/parametric.py +701 -0
  75. twiga/distributions/nn/quantile.py +244 -0
  76. twiga/distributions/nn/residual_conformal.py +273 -0
  77. twiga/forecaster/__init__.py +18 -0
  78. twiga/forecaster/abstract.py +904 -0
  79. twiga/forecaster/base.py +1045 -0
  80. twiga/forecaster/core.py +363 -0
  81. twiga/forecaster/ensemble.py +76 -0
  82. twiga/forecaster/registry.py +99 -0
  83. twiga/forecaster/result.py +362 -0
  84. twiga/forecaster/utils.py +312 -0
  85. twiga/models/__init__.py +94 -0
  86. twiga/models/baseline/__init__.py +20 -0
  87. twiga/models/baseline/context_parrot_model.py +330 -0
  88. twiga/models/baseline/drift_model.py +179 -0
  89. twiga/models/baseline/naive_model.py +222 -0
  90. twiga/models/baseline/seasonal_naive_model.py +285 -0
  91. twiga/models/baseline/window_average_model.py +206 -0
  92. twiga/models/ml/__init__.py +80 -0
  93. twiga/models/ml/catboost_model.py +98 -0
  94. twiga/models/ml/core/__init__.py +0 -0
  95. twiga/models/ml/core/base_regressor.py +234 -0
  96. twiga/models/ml/gausscatboost_model.py +290 -0
  97. twiga/models/ml/lightgbm_model.py +197 -0
  98. twiga/models/ml/lineareg_model.py +74 -0
  99. twiga/models/ml/ngboostexponential_model.py +160 -0
  100. twiga/models/ml/ngboostlognormal_model.py +166 -0
  101. twiga/models/ml/ngboostnormal_model.py +152 -0
  102. twiga/models/ml/prob/__init__.py +0 -0
  103. twiga/models/ml/prob/base_ngboost.py +149 -0
  104. twiga/models/ml/prob/base_quantile.py +318 -0
  105. twiga/models/ml/qrcatboost_model.py +96 -0
  106. twiga/models/ml/qrlightgbm_model.py +71 -0
  107. twiga/models/ml/qrrandomforest_model.py +196 -0
  108. twiga/models/ml/qrxgboost_model.py +85 -0
  109. twiga/models/ml/randomforest_model.py +115 -0
  110. twiga/models/ml/xgboost_model.py +162 -0
  111. twiga/models/nn/__init__.py +140 -0
  112. twiga/models/nn/core/__init__.py +46 -0
  113. twiga/models/nn/core/base.py +376 -0
  114. twiga/models/nn/core/base_arch.py +191 -0
  115. twiga/models/nn/core/base_model.py +541 -0
  116. twiga/models/nn/core/defaults.py +305 -0
  117. twiga/models/nn/core/embedding.py +673 -0
  118. twiga/models/nn/core/linear.py +729 -0
  119. twiga/models/nn/core/scheduler.py +0 -0
  120. twiga/models/nn/mlpf_model.py +187 -0
  121. twiga/models/nn/mlpfbeta_model.py +139 -0
  122. twiga/models/nn/mlpfcrc_model.py +142 -0
  123. twiga/models/nn/mlpffpqr_model.py +207 -0
  124. twiga/models/nn/mlpfgamma_model.py +139 -0
  125. twiga/models/nn/mlpflaplace_model.py +128 -0
  126. twiga/models/nn/mlpflognormal_model.py +132 -0
  127. twiga/models/nn/mlpfnormal_model.py +127 -0
  128. twiga/models/nn/mlpfqr_model.py +381 -0
  129. twiga/models/nn/mlpfstudentt_model.py +135 -0
  130. twiga/models/nn/mlpgaf_model.py +224 -0
  131. twiga/models/nn/mlpgafbeta_model.py +142 -0
  132. twiga/models/nn/mlpgafcrc_model.py +150 -0
  133. twiga/models/nn/mlpgaffpqr_model.py +169 -0
  134. twiga/models/nn/mlpgafgamma_model.py +142 -0
  135. twiga/models/nn/mlpgaflaplace_model.py +132 -0
  136. twiga/models/nn/mlpgaflognormal_model.py +132 -0
  137. twiga/models/nn/mlpgafnormal_model.py +139 -0
  138. twiga/models/nn/mlpgafqr_model.py +166 -0
  139. twiga/models/nn/mlpgafstudentt_model.py +137 -0
  140. twiga/models/nn/mlpgam_model.py +169 -0
  141. twiga/models/nn/mlpgambeta_model.py +138 -0
  142. twiga/models/nn/mlpgamcrc_model.py +136 -0
  143. twiga/models/nn/mlpgamfpqr_model.py +163 -0
  144. twiga/models/nn/mlpgamgamma_model.py +138 -0
  145. twiga/models/nn/mlpgamlaplace_model.py +128 -0
  146. twiga/models/nn/mlpgamlognormal_model.py +128 -0
  147. twiga/models/nn/mlpgamnormal_model.py +134 -0
  148. twiga/models/nn/mlpgamqr_model.py +162 -0
  149. twiga/models/nn/mlpgamstudentt_model.py +133 -0
  150. twiga/models/nn/net/__init__.py +17 -0
  151. twiga/models/nn/net/ganf/__init__.py +0 -0
  152. twiga/models/nn/net/ganf/gates.py +73 -0
  153. twiga/models/nn/net/ganf/groups.py +170 -0
  154. twiga/models/nn/net/ganf/importance.py +829 -0
  155. twiga/models/nn/net/ganf/losses.py +159 -0
  156. twiga/models/nn/net/mlpf.py +414 -0
  157. twiga/models/nn/net/mlpfgaf.py +589 -0
  158. twiga/models/nn/net/mlpfgam.py +468 -0
  159. twiga/models/nn/net/nhits.py +723 -0
  160. twiga/models/nn/net/rnn.py +257 -0
  161. twiga/models/nn/nhits_model.py +306 -0
  162. twiga/models/nn/nhitsbeta_model.py +98 -0
  163. twiga/models/nn/nhitscrc_model.py +122 -0
  164. twiga/models/nn/nhitsgamma_model.py +98 -0
  165. twiga/models/nn/nhitslaplace_model.py +98 -0
  166. twiga/models/nn/nhitslognormal_model.py +98 -0
  167. twiga/models/nn/nhitsnormal_model.py +98 -0
  168. twiga/models/nn/nhitsqr_model.py +239 -0
  169. twiga/models/nn/nhitsstudentt_model.py +101 -0
  170. twiga/models/nn/prob/__init__.py +81 -0
  171. twiga/models/nn/prob/core.py +201 -0
  172. twiga/models/nn/prob/mlpf_crc.py +120 -0
  173. twiga/models/nn/prob/mlpf_parametric.py +456 -0
  174. twiga/models/nn/prob/mlpffpqr.py +144 -0
  175. twiga/models/nn/prob/mlpfqr.py +143 -0
  176. twiga/models/nn/prob/mlpgaf_crc.py +130 -0
  177. twiga/models/nn/prob/mlpgaf_parametric.py +499 -0
  178. twiga/models/nn/prob/mlpgaffpqr.py +162 -0
  179. twiga/models/nn/prob/mlpgafqr.py +162 -0
  180. twiga/models/nn/prob/mlpgam_crc.py +118 -0
  181. twiga/models/nn/prob/mlpgam_parametric.py +453 -0
  182. twiga/models/nn/prob/mlpgamfpqr.py +143 -0
  183. twiga/models/nn/prob/mlpgamqr.py +142 -0
  184. twiga/models/nn/prob/nhits_crc.py +107 -0
  185. twiga/models/nn/prob/nhits_parametric.py +412 -0
  186. twiga/models/nn/prob/nhitsfpqr.py +132 -0
  187. twiga/models/nn/prob/nhitsqr.py +131 -0
  188. twiga/models/nn/prob/rnn_parametric.py +500 -0
  189. twiga/models/nn/prob/rnn_qr.py +242 -0
  190. twiga/models/nn/rnn_model.py +159 -0
  191. twiga/models/nn/rnnbeta_model.py +108 -0
  192. twiga/models/nn/rnngamma_model.py +108 -0
  193. twiga/models/nn/rnnlaplace_model.py +108 -0
  194. twiga/models/nn/rnnlognormal_model.py +108 -0
  195. twiga/models/nn/rnnnormal_model.py +108 -0
  196. twiga/models/nn/rnnqr_model.py +251 -0
  197. twiga/models/nn/rnnstudentt_model.py +108 -0
  198. twiga/pipeline/__init__.py +52 -0
  199. twiga/pipeline/flows.py +441 -0
  200. twiga/pipeline/schemas.py +98 -0
  201. twiga/pipeline/tasks.py +350 -0
  202. twiga/serve/__init__.py +41 -0
  203. twiga/serve/app.py +408 -0
  204. twiga/serve/loader.py +174 -0
  205. twiga/serve/monitor.py +462 -0
  206. twiga/serve/schemas.py +263 -0
  207. twiga/tracking/__init__.py +22 -0
  208. twiga/tracking/tracker.py +314 -0
  209. twiga/utils/__init__.py +7 -0
  210. twiga/utils/data_loading.py +281 -0
  211. twiga/utils/run_forecast.py +757 -0
  212. twiga/utils/utils.py +36 -0
  213. twiga-0.1.0.dist-info/METADATA +467 -0
  214. twiga-0.1.0.dist-info/RECORD +216 -0
  215. twiga-0.1.0.dist-info/WHEEL +4 -0
  216. 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
+ ]
@@ -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
+ ]