arimasel 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.
- arimasel/__init__.py +37 -0
- arimasel/_utils.py +99 -0
- arimasel/cart_arima.py +290 -0
- arimasel/cv.py +130 -0
- arimasel/data/exchange_ng.csv +169 -0
- arimasel/data/gdp_ng.csv +35 -0
- arimasel/data/inflation_ng.csv +289 -0
- arimasel/datasets.py +39 -0
- arimasel/diagnostics.py +146 -0
- arimasel/eda.py +270 -0
- arimasel/forecast.py +151 -0
- arimasel/plotting.py +283 -0
- arimasel/weights.py +113 -0
- arimasel-0.1.0.dist-info/METADATA +147 -0
- arimasel-0.1.0.dist-info/RECORD +18 -0
- arimasel-0.1.0.dist-info/WHEEL +5 -0
- arimasel-0.1.0.dist-info/licenses/LICENSE +674 -0
- arimasel-0.1.0.dist-info/top_level.txt +1 -0
arimasel/__init__.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
"""
|
|
2
|
+
arimasel: Cartesian Product-Based (Seasonal) ARIMA Model Identification and Selection
|
|
3
|
+
======================================================================================
|
|
4
|
+
|
|
5
|
+
A Python port of the R package of the same name. Provides an exhaustive,
|
|
6
|
+
transparent alternative to stepwise automatic ARIMA search: given
|
|
7
|
+
user-supplied index sets P, D, Q (optionally combined with seasonal sets at
|
|
8
|
+
a given period), every candidate (p,d,q)(P,D,Q)[m] model is fit, ranked
|
|
9
|
+
simultaneously by AIC, AICc, BIC, and HQIC, and summarised with Akaike
|
|
10
|
+
weights. Also provides exogenous-regressor support, ensemble forecasting,
|
|
11
|
+
rolling-origin cross-validation, feature-based exploratory data analysis,
|
|
12
|
+
and a feature-guided automatic search (``smart_arima``).
|
|
13
|
+
|
|
14
|
+
Author: Olushina Olawale Awe <olawaleawe@gmail.com>
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
__version__ = "0.1.0"
|
|
18
|
+
|
|
19
|
+
from .cart_arima import CartArimaResult, arima_table, cart_arima
|
|
20
|
+
from .cv import ArimaCVResult, arima_cv
|
|
21
|
+
from .datasets import load_exchange_ng, load_gdp_ng, load_inflation_ng
|
|
22
|
+
from .diagnostics import DiagnoseResult, arima_diagnose, stationarity_test, suggest_d
|
|
23
|
+
from .eda import EdaResult, seasonal_strength, smart_arima, suggest_D, ts_eda, ts_features
|
|
24
|
+
from .forecast import ForecastResult, arima_forecast
|
|
25
|
+
from .weights import arima_weights, cp_sets, hqic
|
|
26
|
+
|
|
27
|
+
__all__ = [
|
|
28
|
+
"__version__",
|
|
29
|
+
"cart_arima", "CartArimaResult", "arima_table",
|
|
30
|
+
"arima_forecast", "ForecastResult",
|
|
31
|
+
"arima_diagnose", "DiagnoseResult",
|
|
32
|
+
"stationarity_test", "suggest_d",
|
|
33
|
+
"arima_cv", "ArimaCVResult",
|
|
34
|
+
"seasonal_strength", "suggest_D", "ts_features", "ts_eda", "EdaResult", "smart_arima",
|
|
35
|
+
"hqic", "cp_sets", "arima_weights",
|
|
36
|
+
"load_gdp_ng", "load_inflation_ng", "load_exchange_ng",
|
|
37
|
+
]
|
arimasel/_utils.py
ADDED
|
@@ -0,0 +1,99 @@
|
|
|
1
|
+
"""Internal helper utilities for arimasel. Not part of the public API."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import warnings
|
|
5
|
+
from typing import Optional, Sequence
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import pandas as pd
|
|
9
|
+
from statsmodels.tsa.arima.model import ARIMA
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _as_series(x) -> pd.Series:
|
|
13
|
+
"""Coerce input to a plain pandas Series with a simple RangeIndex,
|
|
14
|
+
preserving values only (any datetime index is intentionally dropped
|
|
15
|
+
to keep model-fitting behaviour independent of index type -- callers
|
|
16
|
+
that want a real datetime axis can re-attach one on the way out)."""
|
|
17
|
+
if isinstance(x, pd.Series):
|
|
18
|
+
return pd.Series(np.asarray(x.values, dtype=float))
|
|
19
|
+
if isinstance(x, pd.DataFrame):
|
|
20
|
+
if x.shape[1] != 1:
|
|
21
|
+
raise ValueError("DataFrame input must have exactly one column.")
|
|
22
|
+
return pd.Series(np.asarray(x.iloc[:, 0].values, dtype=float))
|
|
23
|
+
arr = np.asarray(x, dtype=float)
|
|
24
|
+
if arr.ndim != 1:
|
|
25
|
+
raise ValueError("Input series must be one-dimensional.")
|
|
26
|
+
return pd.Series(arr)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _model_string(p, d, q, P=None, D=None, Q=None, period=None) -> str:
|
|
30
|
+
base = f"ARIMA({p},{d},{q})"
|
|
31
|
+
if P is not None and D is not None and Q is not None and period is not None:
|
|
32
|
+
base = f"{base}({P},{D},{Q})[{period}]"
|
|
33
|
+
return base
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _safe_fit(x: np.ndarray, order, seasonal_order=None, exog=None, **kwargs):
|
|
37
|
+
"""Fit an ARIMA/SARIMAX model, swallowing convergence warnings and
|
|
38
|
+
returning None (rather than raising) on failure -- candidate models
|
|
39
|
+
in an exhaustive search are *expected* to fail sometimes."""
|
|
40
|
+
try:
|
|
41
|
+
with warnings.catch_warnings():
|
|
42
|
+
warnings.simplefilter("ignore")
|
|
43
|
+
model_kwargs = dict(order=order)
|
|
44
|
+
if seasonal_order is not None:
|
|
45
|
+
model_kwargs["seasonal_order"] = seasonal_order
|
|
46
|
+
if exog is not None:
|
|
47
|
+
model_kwargs["exog"] = exog
|
|
48
|
+
model = ARIMA(x, **model_kwargs)
|
|
49
|
+
fit = model.fit(**kwargs)
|
|
50
|
+
if not np.isfinite(fit.llf):
|
|
51
|
+
return None
|
|
52
|
+
return fit
|
|
53
|
+
except Exception:
|
|
54
|
+
return None
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _extract_metrics(fit) -> dict:
|
|
58
|
+
return {
|
|
59
|
+
"LogLik": fit.llf,
|
|
60
|
+
"AIC": fit.aic,
|
|
61
|
+
"AICc": fit.aicc,
|
|
62
|
+
"BIC": fit.bic,
|
|
63
|
+
"HQIC": fit.hqic,
|
|
64
|
+
"n_params": int(fit.params.shape[0]),
|
|
65
|
+
}
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _vote_table(full_table: pd.DataFrame) -> pd.DataFrame:
|
|
69
|
+
criteria = ["AIC", "AICc", "BIC", "HQIC"]
|
|
70
|
+
winners = {c: full_table.loc[full_table[c].idxmin(), "Model"] for c in criteria}
|
|
71
|
+
counts: dict[str, list[str]] = {}
|
|
72
|
+
for c, m in winners.items():
|
|
73
|
+
counts.setdefault(m, []).append(c)
|
|
74
|
+
rows = [
|
|
75
|
+
{"Model": m, "Votes": len(crits), "Criteria": "/".join(crits)}
|
|
76
|
+
for m, crits in counts.items()
|
|
77
|
+
]
|
|
78
|
+
out = pd.DataFrame(rows).sort_values("Votes", ascending=False).reset_index(drop=True)
|
|
79
|
+
return out
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _validate_seasonal(seasonal: Optional[dict]) -> Optional[dict]:
|
|
83
|
+
if seasonal is None:
|
|
84
|
+
return None
|
|
85
|
+
if "period" not in seasonal or seasonal["period"] is None or seasonal["period"] < 2:
|
|
86
|
+
raise ValueError("seasonal['period'] must be supplied as an integer >= 2.")
|
|
87
|
+
P = sorted(set(seasonal.get("P", range(2))))
|
|
88
|
+
D = sorted(set(seasonal.get("D", range(2))))
|
|
89
|
+
Q = sorted(set(seasonal.get("Q", range(2))))
|
|
90
|
+
if any(v < 0 for v in P + D + Q):
|
|
91
|
+
raise ValueError("Seasonal index sets P, D, Q must be >= 0.")
|
|
92
|
+
return {"P": P, "D": D, "Q": Q, "period": int(seasonal["period"])}
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def _akaike_weights(values: Sequence[float]) -> np.ndarray:
|
|
96
|
+
values = np.asarray(values, dtype=float)
|
|
97
|
+
delta = values - np.nanmin(values)
|
|
98
|
+
raw = np.exp(-delta / 2.0)
|
|
99
|
+
return raw / np.nansum(raw)
|
arimasel/cart_arima.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
1
|
+
"""Cartesian product (seasonal) ARIMA model identification and selection."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import itertools
|
|
5
|
+
import warnings
|
|
6
|
+
from concurrent.futures import ProcessPoolExecutor
|
|
7
|
+
from typing import Optional, Sequence, Union
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
import pandas as pd
|
|
11
|
+
|
|
12
|
+
from ._utils import (
|
|
13
|
+
_akaike_weights,
|
|
14
|
+
_as_series,
|
|
15
|
+
_extract_metrics,
|
|
16
|
+
_model_string,
|
|
17
|
+
_safe_fit,
|
|
18
|
+
_validate_seasonal,
|
|
19
|
+
_vote_table,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
_CRITERIA = ("AIC", "AICc", "BIC", "HQIC")
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class CartArimaResult:
|
|
26
|
+
"""Result of :func:`cart_arima`.
|
|
27
|
+
|
|
28
|
+
Attributes
|
|
29
|
+
----------
|
|
30
|
+
best_model : statsmodels ARIMAResults
|
|
31
|
+
The fitted model for the top-ranked candidate.
|
|
32
|
+
best_model_str : str
|
|
33
|
+
e.g. ``"ARIMA(1,1,1)"`` or ``"ARIMA(1,1,1)(1,0,0)[12]"``.
|
|
34
|
+
criterion : str
|
|
35
|
+
The primary ranking criterion used.
|
|
36
|
+
table : pandas.DataFrame
|
|
37
|
+
The top ``top_n`` candidate models.
|
|
38
|
+
full_table : pandas.DataFrame
|
|
39
|
+
Every converged candidate model, ranked.
|
|
40
|
+
vote_table : pandas.DataFrame
|
|
41
|
+
How many of the four criteria (AIC, AICc, BIC, HQIC) each model wins.
|
|
42
|
+
failed_models : list[str]
|
|
43
|
+
Candidate model strings that failed to converge.
|
|
44
|
+
n_total, n_converged : int
|
|
45
|
+
p_set, d_set, q_set : list[int]
|
|
46
|
+
seasonal : dict or None
|
|
47
|
+
exog : numpy.ndarray or None
|
|
48
|
+
data : pandas.Series
|
|
49
|
+
The original series used for fitting.
|
|
50
|
+
n_obs : int
|
|
51
|
+
"""
|
|
52
|
+
|
|
53
|
+
def __init__(self, **kwargs):
|
|
54
|
+
self.__dict__.update(kwargs)
|
|
55
|
+
self.eda = None # populated by smart_arima()
|
|
56
|
+
|
|
57
|
+
def __repr__(self):
|
|
58
|
+
return self.summary_str()
|
|
59
|
+
|
|
60
|
+
def summary_str(self) -> str:
|
|
61
|
+
lines = []
|
|
62
|
+
lines.append("Cartesian Product ARIMA Model Selection")
|
|
63
|
+
lines.append("=" * 41)
|
|
64
|
+
lines.append(f"Obs (n) : {self.n_obs}")
|
|
65
|
+
if self.seasonal is not None:
|
|
66
|
+
s = self.seasonal
|
|
67
|
+
lines.append(
|
|
68
|
+
f"Seasonal : period={s['period']}, "
|
|
69
|
+
f"P={s['P']}, D={s['D']}, Q={s['Q']}"
|
|
70
|
+
)
|
|
71
|
+
if self.exog is not None:
|
|
72
|
+
lines.append(f"exog : {self.exog.shape[1]} external regressor column(s)")
|
|
73
|
+
lines.append(
|
|
74
|
+
f"Candidates: {self.n_total} | Converged: {self.n_converged} "
|
|
75
|
+
f"| Failed: {self.n_total - self.n_converged}"
|
|
76
|
+
)
|
|
77
|
+
lines.append(f"Criterion : {self.criterion}")
|
|
78
|
+
lines.append(f"Best model: {self.best_model_str}\n")
|
|
79
|
+
lines.append(f"Top {len(self.table)} models (ranked by {self.criterion}):")
|
|
80
|
+
cols = ["Rank", "Model", "AIC", "AICc", "BIC", "HQIC", "Delta", "Weight"]
|
|
81
|
+
lines.append(self.table[cols].to_string(index=False))
|
|
82
|
+
if self.failed_models:
|
|
83
|
+
lines.append(
|
|
84
|
+
f"\nFailed to converge ({len(self.failed_models)}): "
|
|
85
|
+
+ ", ".join(self.failed_models)
|
|
86
|
+
)
|
|
87
|
+
return "\n".join(lines)
|
|
88
|
+
|
|
89
|
+
def summary(self) -> None:
|
|
90
|
+
print(self.summary_str())
|
|
91
|
+
|
|
92
|
+
@property
|
|
93
|
+
def fittedvalues(self) -> np.ndarray:
|
|
94
|
+
return np.asarray(self.best_model.fittedvalues)
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def resid(self) -> np.ndarray:
|
|
98
|
+
return np.asarray(self.best_model.resid)
|
|
99
|
+
|
|
100
|
+
def plot(self, kind: str = "criteria", criterion: Optional[str] = None,
|
|
101
|
+
top_n: int = 10, ax=None, **kwargs):
|
|
102
|
+
from .plotting import plot_cart_arima
|
|
103
|
+
return plot_cart_arima(self, kind=kind, criterion=criterion,
|
|
104
|
+
top_n=top_n, ax=ax, **kwargs)
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def _fit_one(args):
|
|
108
|
+
(x_vals, p, d, q, P, D, Q, period, exog, fit_kwargs) = args
|
|
109
|
+
seasonal_order = (P, D, Q, period) if period is not None else None
|
|
110
|
+
model_str = _model_string(p, d, q, P, D, Q, period)
|
|
111
|
+
fit = _safe_fit(x_vals, order=(p, d, q), seasonal_order=seasonal_order,
|
|
112
|
+
exog=exog, **fit_kwargs)
|
|
113
|
+
if fit is None:
|
|
114
|
+
return {"row": None, "fit": None, "model_str": model_str}
|
|
115
|
+
m = _extract_metrics(fit)
|
|
116
|
+
row = {
|
|
117
|
+
"Model": model_str, "p": p, "d": d, "q": q,
|
|
118
|
+
"P": P, "D": D, "Q": Q,
|
|
119
|
+
"LogLik": round(m["LogLik"], 3), "AIC": round(m["AIC"], 3),
|
|
120
|
+
"AICc": round(m["AICc"], 3), "BIC": round(m["BIC"], 3),
|
|
121
|
+
"HQIC": round(m["HQIC"], 3),
|
|
122
|
+
}
|
|
123
|
+
return {"row": row, "fit": fit, "model_str": model_str}
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def cart_arima(
|
|
127
|
+
x,
|
|
128
|
+
p_set: Sequence[int] = range(3),
|
|
129
|
+
d_set: Sequence[int] = range(2),
|
|
130
|
+
q_set: Sequence[int] = range(3),
|
|
131
|
+
seasonal: Optional[dict] = None,
|
|
132
|
+
exog: Optional[Union[np.ndarray, pd.DataFrame]] = None,
|
|
133
|
+
criterion: str = "AIC",
|
|
134
|
+
top_n: int = 10,
|
|
135
|
+
n_jobs: int = 1,
|
|
136
|
+
**fit_kwargs,
|
|
137
|
+
) -> CartArimaResult:
|
|
138
|
+
"""Exhaustive Cartesian-product (seasonal) ARIMA search.
|
|
139
|
+
|
|
140
|
+
Fits every candidate ``(p, d, q)`` -- optionally combined with a seasonal
|
|
141
|
+
``(P, D, Q)[period]`` -- in the Cartesian product of the supplied index
|
|
142
|
+
sets, ranks all converged models simultaneously by AIC, AICc, BIC, and
|
|
143
|
+
HQIC, and returns the best model together with the full comparison
|
|
144
|
+
table and Akaike weights.
|
|
145
|
+
|
|
146
|
+
Parameters
|
|
147
|
+
----------
|
|
148
|
+
x : array-like, pandas.Series, or pandas.DataFrame (one column)
|
|
149
|
+
The time series to model. Must have at least 10 observations.
|
|
150
|
+
p_set, d_set, q_set : sequence of int
|
|
151
|
+
Non-negative integer index sets for the non-seasonal AR, differencing,
|
|
152
|
+
and MA orders. Defaults: ``range(3)``, ``range(2)``, ``range(3)``.
|
|
153
|
+
seasonal : dict, optional
|
|
154
|
+
``{"P": <seq>, "D": <seq>, "Q": <seq>, "period": <int>}``. ``P``,
|
|
155
|
+
``D``, ``Q`` each default to ``range(2)`` if omitted; ``period`` is
|
|
156
|
+
required. When supplied, every combination of ``(p,d,q)`` and
|
|
157
|
+
``(P,D,Q)`` at the given period is fitted.
|
|
158
|
+
exog : array-like, optional
|
|
159
|
+
External regressors (regression with ARIMA errors), same length as ``x``.
|
|
160
|
+
criterion : {"AIC", "AICc", "BIC", "HQIC"}
|
|
161
|
+
Primary ranking criterion.
|
|
162
|
+
top_n : int
|
|
163
|
+
Number of rows to keep in the abbreviated ``.table``; the full,
|
|
164
|
+
ranked table of all converged models is always in ``.full_table``.
|
|
165
|
+
n_jobs : int
|
|
166
|
+
Number of worker processes for fitting candidates in parallel.
|
|
167
|
+
``1`` (default) fits serially in-process.
|
|
168
|
+
**fit_kwargs
|
|
169
|
+
Extra keyword arguments passed to ``ARIMAResults.fit()``, e.g.
|
|
170
|
+
``method="statespace"``.
|
|
171
|
+
|
|
172
|
+
Returns
|
|
173
|
+
-------
|
|
174
|
+
CartArimaResult
|
|
175
|
+
"""
|
|
176
|
+
if criterion not in _CRITERIA:
|
|
177
|
+
raise ValueError(f"criterion must be one of {_CRITERIA}")
|
|
178
|
+
|
|
179
|
+
x_series = _as_series(x)
|
|
180
|
+
x_vals = x_series.to_numpy(dtype=float)
|
|
181
|
+
n = len(x_vals)
|
|
182
|
+
if n < 10:
|
|
183
|
+
raise ValueError("Time series must have at least 10 observations.")
|
|
184
|
+
if n < 20:
|
|
185
|
+
warnings.warn("Time series has fewer than 20 observations; results may be unreliable.")
|
|
186
|
+
if np.isnan(x_vals).any():
|
|
187
|
+
warnings.warn("NaN values found in x. Replacing with the series mean.")
|
|
188
|
+
x_vals = np.where(np.isnan(x_vals), np.nanmean(x_vals), x_vals)
|
|
189
|
+
|
|
190
|
+
p_set = sorted(set(int(v) for v in p_set))
|
|
191
|
+
d_set = sorted(set(int(v) for v in d_set))
|
|
192
|
+
q_set = sorted(set(int(v) for v in q_set))
|
|
193
|
+
if any(v < 0 for v in p_set + d_set + q_set):
|
|
194
|
+
raise ValueError("p_set, d_set, q_set must contain only non-negative integers.")
|
|
195
|
+
|
|
196
|
+
seasonal = _validate_seasonal(seasonal)
|
|
197
|
+
|
|
198
|
+
exog_arr = None
|
|
199
|
+
if exog is not None:
|
|
200
|
+
exog_arr = np.asarray(exog, dtype=float)
|
|
201
|
+
if exog_arr.ndim == 1:
|
|
202
|
+
exog_arr = exog_arr.reshape(-1, 1)
|
|
203
|
+
if exog_arr.shape[0] != n:
|
|
204
|
+
raise ValueError("exog must have the same number of rows as len(x).")
|
|
205
|
+
|
|
206
|
+
if seasonal is None:
|
|
207
|
+
combos = [(p, d, q, None, None, None, None)
|
|
208
|
+
for p, d, q in itertools.product(p_set, d_set, q_set)]
|
|
209
|
+
else:
|
|
210
|
+
combos = [
|
|
211
|
+
(p, d, q, P, D, Q, seasonal["period"])
|
|
212
|
+
for p, d, q, P, D, Q in itertools.product(
|
|
213
|
+
p_set, d_set, q_set, seasonal["P"], seasonal["D"], seasonal["Q"]
|
|
214
|
+
)
|
|
215
|
+
]
|
|
216
|
+
n_total = len(combos)
|
|
217
|
+
if n_total == 0:
|
|
218
|
+
raise ValueError("No candidate models generated; check p_set/d_set/q_set.")
|
|
219
|
+
|
|
220
|
+
tasks = [(x_vals, p, d, q, P, D, Q, period, exog_arr, fit_kwargs)
|
|
221
|
+
for (p, d, q, P, D, Q, period) in combos]
|
|
222
|
+
|
|
223
|
+
if n_jobs and n_jobs > 1:
|
|
224
|
+
with ProcessPoolExecutor(max_workers=n_jobs) as ex:
|
|
225
|
+
results = list(ex.map(_fit_one, tasks))
|
|
226
|
+
else:
|
|
227
|
+
results = [_fit_one(t) for t in tasks]
|
|
228
|
+
|
|
229
|
+
rows, fits, failed = [], [], []
|
|
230
|
+
for r in results:
|
|
231
|
+
if r["row"] is None:
|
|
232
|
+
failed.append(r["model_str"])
|
|
233
|
+
else:
|
|
234
|
+
rows.append(r["row"])
|
|
235
|
+
fits.append(r["fit"])
|
|
236
|
+
|
|
237
|
+
if not rows:
|
|
238
|
+
raise RuntimeError("All candidate (S)ARIMA models failed to converge.")
|
|
239
|
+
|
|
240
|
+
full_table = pd.DataFrame(rows)
|
|
241
|
+
order_idx = np.argsort(full_table[criterion].to_numpy())
|
|
242
|
+
full_table = full_table.iloc[order_idx].reset_index(drop=True)
|
|
243
|
+
fits = [fits[i] for i in order_idx]
|
|
244
|
+
|
|
245
|
+
full_table["Rank"] = np.arange(1, len(full_table) + 1)
|
|
246
|
+
full_table["Delta"] = (full_table[criterion] - full_table[criterion].min()).round(3)
|
|
247
|
+
full_table["Weight"] = _akaike_weights(full_table[criterion].to_numpy()).round(4)
|
|
248
|
+
|
|
249
|
+
vote_table = _vote_table(full_table)
|
|
250
|
+
|
|
251
|
+
best_model = fits[0]
|
|
252
|
+
best_model_str = full_table.loc[0, "Model"]
|
|
253
|
+
|
|
254
|
+
top_n_use = min(int(top_n), len(full_table))
|
|
255
|
+
table = full_table.iloc[:top_n_use].reset_index(drop=True)
|
|
256
|
+
|
|
257
|
+
return CartArimaResult(
|
|
258
|
+
best_model=best_model,
|
|
259
|
+
best_model_str=best_model_str,
|
|
260
|
+
criterion=criterion,
|
|
261
|
+
table=table,
|
|
262
|
+
full_table=full_table,
|
|
263
|
+
vote_table=vote_table,
|
|
264
|
+
failed_models=failed,
|
|
265
|
+
n_total=n_total,
|
|
266
|
+
n_converged=len(full_table),
|
|
267
|
+
p_set=p_set, d_set=d_set, q_set=q_set,
|
|
268
|
+
seasonal=seasonal,
|
|
269
|
+
exog=exog_arr,
|
|
270
|
+
data=x_series,
|
|
271
|
+
n_obs=n,
|
|
272
|
+
_fits=fits,
|
|
273
|
+
)
|
|
274
|
+
|
|
275
|
+
|
|
276
|
+
def arima_table(result: CartArimaResult, criterion: Optional[str] = None,
|
|
277
|
+
top_n: Optional[int] = None) -> pd.DataFrame:
|
|
278
|
+
"""Re-rank and/or truncate the comparison table stored in a
|
|
279
|
+
:class:`CartArimaResult`, without refitting anything."""
|
|
280
|
+
ft = result.full_table.copy()
|
|
281
|
+
if criterion is not None:
|
|
282
|
+
if criterion not in _CRITERIA:
|
|
283
|
+
raise ValueError(f"criterion must be one of {_CRITERIA}")
|
|
284
|
+
ft = ft.sort_values(criterion).reset_index(drop=True)
|
|
285
|
+
ft["Rank"] = np.arange(1, len(ft) + 1)
|
|
286
|
+
ft["Delta"] = (ft[criterion] - ft[criterion].min()).round(3)
|
|
287
|
+
ft["Weight"] = _akaike_weights(ft[criterion].to_numpy()).round(4)
|
|
288
|
+
if top_n is not None:
|
|
289
|
+
ft = ft.iloc[: int(top_n)].reset_index(drop=True)
|
|
290
|
+
return ft
|
arimasel/cv.py
ADDED
|
@@ -0,0 +1,130 @@
|
|
|
1
|
+
"""Rolling-origin (expanding window) cross-validation."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from typing import Optional
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
|
|
9
|
+
from ._utils import _safe_fit
|
|
10
|
+
from .cart_arima import CartArimaResult
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class ArimaCVResult:
|
|
14
|
+
"""Result of :func:`arima_cv`."""
|
|
15
|
+
|
|
16
|
+
def __init__(self, **kwargs):
|
|
17
|
+
self.__dict__.update(kwargs)
|
|
18
|
+
|
|
19
|
+
def __repr__(self):
|
|
20
|
+
lines = ["Rolling-Origin Cross-Validation", "=" * 32]
|
|
21
|
+
order_str = f"ARIMA{self.order}"
|
|
22
|
+
if self.seasonal_order is not None:
|
|
23
|
+
order_str += f"{self.seasonal_order[:3]}[{self.seasonal_order[3]}]"
|
|
24
|
+
lines.append(f"Model order : {order_str}")
|
|
25
|
+
lines.append(f"Origins : {self.n_folds} "
|
|
26
|
+
f"(initial={self.initial}, step={self.step}, h={self.h})\n")
|
|
27
|
+
lines.append("Out-of-sample accuracy by horizon:")
|
|
28
|
+
lines.append(self.accuracy.to_string(index=False))
|
|
29
|
+
return "\n".join(lines)
|
|
30
|
+
|
|
31
|
+
def plot(self, kind: str = "rmse", ax=None, **kwargs):
|
|
32
|
+
from .plotting import plot_cv
|
|
33
|
+
return plot_cv(self, kind=kind, ax=ax, **kwargs)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def arima_cv(
|
|
37
|
+
result: CartArimaResult,
|
|
38
|
+
h: int = 1,
|
|
39
|
+
initial: Optional[int] = None,
|
|
40
|
+
step: int = 1,
|
|
41
|
+
) -> ArimaCVResult:
|
|
42
|
+
"""Rolling-origin (expanding window) cross-validation of the model order
|
|
43
|
+
selected by :func:`cart_arima`: refit the same order at successive
|
|
44
|
+
origins, forecast ``h`` steps ahead, and record out-of-sample error.
|
|
45
|
+
|
|
46
|
+
Parameters
|
|
47
|
+
----------
|
|
48
|
+
result : CartArimaResult
|
|
49
|
+
h : int
|
|
50
|
+
Forecast horizon evaluated at each origin.
|
|
51
|
+
initial : int, optional
|
|
52
|
+
Size of the first training window. Defaults to
|
|
53
|
+
``max(20, floor(0.5 * n_obs))``.
|
|
54
|
+
step : int
|
|
55
|
+
Number of observations the training window expands by between origins.
|
|
56
|
+
|
|
57
|
+
Returns
|
|
58
|
+
-------
|
|
59
|
+
ArimaCVResult
|
|
60
|
+
"""
|
|
61
|
+
h = int(h)
|
|
62
|
+
step = int(step)
|
|
63
|
+
if h < 1:
|
|
64
|
+
raise ValueError("h must be a positive integer.")
|
|
65
|
+
if step < 1:
|
|
66
|
+
raise ValueError("step must be a positive integer.")
|
|
67
|
+
|
|
68
|
+
x = result.data.to_numpy(dtype=float)
|
|
69
|
+
n = len(x)
|
|
70
|
+
|
|
71
|
+
fit = result.best_model
|
|
72
|
+
order = fit.model.order
|
|
73
|
+
seasonal_order = None
|
|
74
|
+
if result.seasonal is not None:
|
|
75
|
+
seasonal_order = fit.model.seasonal_order
|
|
76
|
+
|
|
77
|
+
if initial is None:
|
|
78
|
+
initial = max(20, n // 2)
|
|
79
|
+
initial = int(initial)
|
|
80
|
+
if initial >= n - h:
|
|
81
|
+
raise ValueError("initial leaves no observations for out-of-sample evaluation.")
|
|
82
|
+
|
|
83
|
+
origins = list(range(initial, n - h + 1, step))
|
|
84
|
+
if not origins:
|
|
85
|
+
raise ValueError("No valid CV origins with the given initial/h/step.")
|
|
86
|
+
|
|
87
|
+
exog = result.exog
|
|
88
|
+
rows = []
|
|
89
|
+
for orig in origins:
|
|
90
|
+
train = x[:orig]
|
|
91
|
+
actual = x[orig: orig + h]
|
|
92
|
+
exog_tr = exog[:orig] if exog is not None else None
|
|
93
|
+
exog_te = exog[orig: orig + h] if exog is not None else None
|
|
94
|
+
|
|
95
|
+
fit_o = _safe_fit(train, order=order, seasonal_order=seasonal_order, exog=exog_tr)
|
|
96
|
+
if fit_o is None:
|
|
97
|
+
fc = np.full(h, np.nan)
|
|
98
|
+
else:
|
|
99
|
+
try:
|
|
100
|
+
fc = np.asarray(fit_o.get_forecast(steps=h, exog=exog_te).predicted_mean)
|
|
101
|
+
except Exception:
|
|
102
|
+
fc = np.full(h, np.nan)
|
|
103
|
+
|
|
104
|
+
for j in range(h):
|
|
105
|
+
rows.append({"Origin": orig, "Horizon": j + 1,
|
|
106
|
+
"Actual": actual[j], "Forecast": fc[j],
|
|
107
|
+
"Error": actual[j] - fc[j]})
|
|
108
|
+
|
|
109
|
+
errors = pd.DataFrame(rows)
|
|
110
|
+
|
|
111
|
+
acc_rows = []
|
|
112
|
+
for horizon, grp in errors.groupby("Horizon"):
|
|
113
|
+
e = grp["Error"].to_numpy()
|
|
114
|
+
a = grp["Actual"].to_numpy()
|
|
115
|
+
finite = np.isfinite(e)
|
|
116
|
+
e, a = e[finite], a[finite]
|
|
117
|
+
acc_rows.append({
|
|
118
|
+
"Horizon": horizon,
|
|
119
|
+
"RMSE": round(float(np.sqrt(np.mean(e ** 2))), 4) if len(e) else np.nan,
|
|
120
|
+
"MAE": round(float(np.mean(np.abs(e))), 4) if len(e) else np.nan,
|
|
121
|
+
"MAPE": round(float(np.mean(np.abs(e / a))) * 100, 4) if len(e) else np.nan,
|
|
122
|
+
"N": int(len(e)),
|
|
123
|
+
})
|
|
124
|
+
accuracy = pd.DataFrame(acc_rows).sort_values("Horizon").reset_index(drop=True)
|
|
125
|
+
|
|
126
|
+
return ArimaCVResult(
|
|
127
|
+
errors=errors, accuracy=accuracy,
|
|
128
|
+
order=order, seasonal_order=seasonal_order,
|
|
129
|
+
n_folds=len(origins), h=h, initial=initial, step=step,
|
|
130
|
+
)
|