Test TimesFM 3 on agriculture and forestry time series #2

Draft implementation of a Time series forecasting with Claude Opus 5.

Submission as a merge request to test the issue workflow for future contributors to our modelling meeting.

Merge request reports

Loading
+0 −0

Empty file added.

+0 −0

Empty file added.

+0 −0

Empty file added.

+655 −0
Changes for events/20260917/src/backtest.py: 655 added lines, 0 removed lines.
Original line number Diff line number Diff line
"""Rolling-origin backtest: TimesFM 3.0 against econometric baselines.

Every model sees exactly the same information set at every origin: all
observations up to and including the origin month, nothing after. That rule is
enforced in one place, `iter_origins` + the `history` slice in `run_split`, so
there is a single line to audit for look-ahead bias.

    python src/backtest.py --models rw,snaive,drift,ets,sarima,var --split test
    python src/backtest.py --models all --split test
    python src/backtest.py --models all --split stress
    python src/backtest.py --models all --split honest

Adding a model is one function with the signature

    fit_predict(history, horizon, warm) -> (point[H], quantiles[H, 9], warm_out)

registered in UNIVARIATE_MODELS. `warm` carries the previous origin's fitted
parameters so that iterative models can warm-start; return None if unused.
"""

from __future__ import annotations

import argparse
import contextlib
import os
import sys
import time
import warnings
from pathlib import Path

# Hundreds of small SARIMAX fits oversubscribe the BLAS thread pool and end up
# slower than single-threaded. Must be set before numpy is imported.
for _v in ("OMP_NUM_THREADS", "OPENBLAS_NUM_THREADS", "MKL_NUM_THREADS",
           "NUMEXPR_NUM_THREADS", "VECLIB_MAXIMUM_THREADS"):
    os.environ.setdefault(_v, "1")

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from scipy import stats

sys.path.insert(0, str(Path(__file__).resolve().parent))
from metrics import (  # noqa: E402
    QUANTILE_COLS, QUANTILE_LEVELS, metrics_by_horizon, metrics_overall,
    seasonal_mae_scale,
)

warnings.filterwarnings("ignore")  # statsmodels convergence noise, by the hundred


@contextlib.contextmanager
def quiet():
    """Silence warnings inside a fit.

    A module-level filterwarnings is not enough: statsmodels re-arms the warning
    filter inside its own fit routines, so hundreds of ConvergenceWarnings end up
    on stdout and bury the progress output. Non-convergence is handled explicitly
    below by refusing to warm-start from a failed fit, so suppressing the message
    does not suppress the problem.
    """
    with warnings.catch_warnings():
        warnings.simplefilter("ignore")
        yield

HERE = Path(__file__).resolve().parent.parent
PANEL = HERE / "data" / "processed" / "prices_monthly.csv"
OUTPUT = HERE / "output"

SEASON = 12
Z = np.array([stats.norm.ppf(q) for q in QUANTILE_LEVELS])

# Named origins for the regime-shift stress set. Each is reported on its own:
# averaging a bark-beetle collapse together with 35 quiet months hides exactly
# the behaviour we want to see.
STRESS_ORIGINS = {
    "2018-06-01": "bark-beetle calamity onset",
    "2020-03-01": "COVID-19",
    "2021-09-01": "construction timber surge",
    "2022-02-01": "Russian invasion of Ukraine",
    "2022-08-01": "energy price peak",
    "2023-06-01": "disinflation",
}


# --------------------------------------------------------------------------- #
# Splits
# --------------------------------------------------------------------------- #
def iter_origins(dates: pd.DatetimeIndex, split: str, horizon: int) -> list:
    """Forecast origins for a split.

    Every date is derived from the last observation T, so the design survives a
    data update without anyone editing a constant.

      dev     T-95 .. T-48   48 origins, for tuning
      test    T-47 .. T-12   36 origins, reported once
      stress  six named origins, reported individually
      honest  T-11 .. T-h    after the TimesFM 3.0 pretraining cutoff
    """
    dates = pd.DatetimeIndex(sorted(dates))
    n = len(dates)
    if split == "dev":
        lo, hi = n - 96, n - 48
    elif split == "test":
        lo, hi = n - 48, n - horizon
    elif split == "honest":
        lo, hi = n - 12, n - horizon
    elif split == "stress":
        wanted = pd.DatetimeIndex([pd.Timestamp(d) for d in STRESS_ORIGINS])
        return [d for d in wanted if d in set(dates) and d <= dates[-1] - pd.DateOffset(months=horizon)]
    else:
        raise ValueError(f"unknown split {split!r}")

    lo = max(lo, SEASON * 4)  # need enough history to fit a seasonal model
    if hi <= lo:
        hint = (" The honest split only has the last 12 months of origins, so it "
                "needs a short horizon: --split honest --horizon 6."
                if split == "honest" else "")
        raise SystemExit(
            f"Split {split!r} is empty: {n} observations is too short for "
            f"horizon {horizon}.{hint}"
        )
    return list(dates[lo:hi])


# --------------------------------------------------------------------------- #
# Baselines: point forecast plus empirical quantiles from in-sample errors
# --------------------------------------------------------------------------- #
def _empirical_quantiles(point: np.ndarray, errors_by_h: list) -> np.ndarray:
    """Quantiles from the in-sample distribution of h-step forecast errors.

    Gives the naive baselines an honest predictive distribution, so pinball loss
    and interval coverage are defined for them too — without that, TimesFM would
    win the distributional comparison by walkover.
    """
    out = np.full((len(point), len(QUANTILE_LEVELS)), np.nan)
    for i, errs in enumerate(errors_by_h):
        errs = np.asarray(errs, float)
        errs = errs[np.isfinite(errs)]
        if len(errs) >= 10:
            out[i] = point[i] + np.quantile(errs, QUANTILE_LEVELS)
        elif len(errs) >= 2:  # too few for empirical quantiles, fall back to normal
            out[i] = point[i] + Z * errs.std(ddof=1)
    return out


def forecast_rw(history: pd.Series, horizon: int, warm=None):
    y = history.to_numpy(float)
    point = np.repeat(y[-1], horizon)
    errors = [y[h:] - y[:-h] if len(y) > h else np.array([]) for h in range(1, horizon + 1)]
    return point, _empirical_quantiles(point, errors), None


def forecast_snaive(history: pd.Series, horizon: int, warm=None):
    y = history.to_numpy(float)
    if len(y) < SEASON:
        return forecast_rw(history, horizon)
    # Reference value: the same month of the year, stepping back in whole years
    # until we land inside the history. Index -1 is the origin month itself.
    point = np.array(
        [y[-1 + h - SEASON * int(np.ceil(h / SEASON))] for h in range(1, horizon + 1)],
        float,
    )
    errors = []
    for h in range(1, horizon + 1):
        lag = SEASON * int(np.ceil(h / SEASON))
        errors.append(y[lag:] - y[:-lag] if len(y) > lag else np.array([]))
    return point, _empirical_quantiles(point, errors), None


def forecast_drift(history: pd.Series, horizon: int, warm=None):
    y = history.to_numpy(float)
    slope = (y[-1] - y[0]) / (len(y) - 1) if len(y) > 1 else 0.0
    point = y[-1] + slope * np.arange(1, horizon + 1)
    errors = []
    for h in range(1, horizon + 1):
        errors.append(y[h:] - (y[:-h] + slope * h) if len(y) > h else np.array([]))
    return point, _empirical_quantiles(point, errors), None


# --------------------------------------------------------------------------- #
# Econometric models. Fitted on logs: prices are positive and multiplicative,
# and a monotone transform maps quantiles to quantiles, so exp() of a log-space
# quantile is the level quantile. No smearing correction — these are median
# forecasts, which is what MASE and pinball loss ask for.
# --------------------------------------------------------------------------- #
def forecast_sarima(history: pd.Series, horizon: int, warm=None):
    from statsmodels.tsa.statespace.sarimax import SARIMAX

    ylog = np.log(history.to_numpy(float))
    model = SARIMAX(
        ylog, order=(1, 1, 1), seasonal_order=(0, 1, 1, SEASON),
        trend="n", enforce_stationarity=False, enforce_invertibility=False,
    )
    # Consecutive origins differ by one observation, so the previous origin's
    # estimates are an excellent starting point. Same specification, same
    # likelihood, just fewer optimiser iterations — roughly a 3x speedup over a
    # 36-origin backtest, with identical optima up to optimiser tolerance.
    with quiet():
        try:
            res = (model.fit(disp=False, start_params=warm) if warm is not None
                   else model.fit(disp=False))
        except Exception:
            res = model.fit(disp=False)

    fc = res.get_forecast(steps=horizon)
    mean, se = np.asarray(fc.predicted_mean, float), np.asarray(fc.se_mean, float)
    quant = np.exp(mean[:, None] + Z[None, :] * se[:, None])

    # Never warm-start the next origin from a fit that did not converge: that
    # would carry a bad optimum forward through the whole backtest instead of
    # letting the next origin retry from the default starting values.
    converged = bool(getattr(res, "mle_retvals", {}).get("converged", False))
    return np.exp(mean), quant, (np.asarray(res.params, float) if converged else None)


def forecast_ets(history: pd.Series, horizon: int, warm=None):
    from statsmodels.tsa.holtwinters import ExponentialSmoothing

    ylog = np.log(history.to_numpy(float))
    seasonal = "add" if len(ylog) >= 2 * SEASON + 4 else None
    with quiet():
        res = ExponentialSmoothing(
            ylog, trend="add", damped_trend=True, seasonal=seasonal,
            seasonal_periods=SEASON if seasonal else None,
            initialization_method="estimated",
        ).fit(optimized=True)
    mean = np.asarray(res.forecast(horizon), float)
    # Holt-Winters has no closed-form predictive variance here. Approximate it by
    # the residual sd growing with sqrt(h) — a random-walk-ish assumption that is
    # generous at long horizons. Flagged rather than hidden.
    sigma = float(np.std(res.resid[np.isfinite(res.resid)], ddof=1))
    se = sigma * np.sqrt(np.arange(1, horizon + 1))
    quant = np.exp(mean[:, None] + Z[None, :] * se[:, None])
    return np.exp(mean), quant, None


UNIVARIATE_MODELS = {
    "rw": forecast_rw,
    "snaive": forecast_snaive,
    "drift": forecast_drift,
    "ets": forecast_ets,
    "sarima": forecast_sarima,
}


# --------------------------------------------------------------------------- #
# VAR: the econometric multivariate comparator for timesfm_mv
# --------------------------------------------------------------------------- #
def forecast_var(history: pd.DataFrame, horizon: int) -> dict:
    """VAR on log first differences, lag order by AIC. Returns {series_id: (point, quant)}."""
    from statsmodels.tsa.api import VAR

    logs = np.log(history.to_numpy(float))
    dlog = np.diff(logs, axis=0)
    k = dlog.shape[1]
    maxlags = min(12, max(1, (len(dlog) - k - 2) // (k + 1)))
    with quiet():
        res = VAR(dlog).fit(maxlags=maxlags, ic="aic")

    steps_dlog = res.forecast(dlog[-res.k_ar:], horizon)          # (H, k)
    log_point = logs[-1] + np.cumsum(steps_dlog, axis=0)

    # Forecast error variance of the *cumulated* differences. res.mse(h) is the
    # covariance of the h-step error on the differences; summing the diagonals
    # ignores the cross-horizon correlation, so this understates the true level
    # uncertainty. Adequate for ranking coverage, not for a risk statement.
    mse = res.mse(horizon)                                        # (H, k, k)
    var_cum = np.cumsum(np.array([np.diag(m) for m in mse]), axis=0)
    se = np.sqrt(var_cum)                                         # (H, k)

    out = {}
    for j, sid in enumerate(history.columns):
        mean = log_point[:, j]
        quant = np.exp(mean[:, None] + Z[None, :] * se[:, j][:, None])
        out[sid] = (np.exp(mean), quant)
    return out


# --------------------------------------------------------------------------- #
# TimesFM 3.0
# --------------------------------------------------------------------------- #
_FORECASTER = None


def get_forecaster(device: str, checkpoint: str, batch_size: int):
    """Load the checkpoint once and keep it. ~1.3 GB on the first call."""
    global _FORECASTER
    if _FORECASTER is None:
        try:
            from timesfm3 import TimesFM3Forecaster
        except ImportError:
            raise SystemExit(
                "timesfm is not installed. `pip install 'timesfm[torch]==3.0.2'`, "
                "or run with --models rw,snaive,drift,ets,sarima,var to check the "
                "harness without it."
            )
        print(f"Loading {checkpoint} on {device} ...")
        _FORECASTER = TimesFM3Forecaster.from_pretrained(
            checkpoint, device=device, per_core_batch_size=batch_size
        )
    return _FORECASTER


def run_timesfm_univariate(wide: pd.DataFrame, origins, horizon: int,
                           context_length: int, forecaster) -> list:
    """One context per (origin, series). All of them go through in one batch."""
    contexts, keys = [], []
    for origin in origins:
        hist = wide.loc[:origin]
        for sid in wide.columns:
            y = hist[sid].dropna().to_numpy(np.float32)
            if len(y) < SEASON * 2:
                continue
            contexts.append(y[-context_length:])
            keys.append((origin, sid))

    print(f"  timesfm_uni: {len(contexts)} series-origins in one batch")
    outputs = list(forecaster.predict_batch(
        contexts=contexts, horizon=horizon, return_quantiles=True,
    ))
    return [(k, np.asarray(o.forecast, float), np.asarray(o.quantiles, float))
            for k, o in zip(keys, outputs)]


def run_timesfm_multivariate(wide: pd.DataFrame, origins, horizon: int,
                             context_length: int, forecaster) -> list:
    """All channels of a group forecast jointly — the headline feature of 3.0.

    Context shape is (num_variates, context_length); the output is (k, H) and
    (k, H, 9). Comparing this against timesfm_uni isolates what the multivariate
    attention actually buys on these series.
    """
    contexts, keys = [], []
    for origin in origins:
        hist = wide.loc[:origin].dropna()
        if len(hist) < SEASON * 2:
            continue
        block = hist.to_numpy(np.float32).T[:, -context_length:]  # (k, L)
        contexts.append(block)
        keys.append(origin)

    print(f"  timesfm_mv: {len(contexts)} origins x {wide.shape[1]} variates")
    outputs = list(forecaster.predict_batch(
        contexts=contexts, horizon=horizon, return_quantiles=True,
    ))

    results = []
    for origin, out in zip(keys, outputs):
        point = np.asarray(out.forecast, float)       # (k, H)
        quant = np.asarray(out.quantiles, float)      # (k, H, 9)
        for j, sid in enumerate(wide.columns):
            results.append(((origin, sid), point[j], quant[j]))
    return results


# --------------------------------------------------------------------------- #
# The harness
# --------------------------------------------------------------------------- #
def collect_rows(key, model, point, quant, wide, horizon) -> list:
    """Attach realised values to a forecast and flatten to tidy rows."""
    origin, sid = key
    rows = []
    for h in range(1, horizon + 1):
        target = origin + pd.DateOffset(months=h)
        if target not in wide.index:
            continue
        row = {
            "origin": origin, "series_id": sid, "model": model, "horizon": h,
            "target_date": target, "y_true": float(wide.at[target, sid]),
            "y_pred": float(point[h - 1]),
        }
        if quant is not None and np.ndim(quant) == 2 and quant.shape[0] >= h:
            row.update(dict(zip(QUANTILE_COLS, quant[h - 1].astype(float))))
        rows.append(row)
    return rows


def run_split(panel: pd.DataFrame, models: list, split: str, horizon: int,
              context_length: int, device: str, checkpoint: str,
              batch_size: int):
    all_rows, scales = [], {}

    for group, gdf in panel.groupby("group"):
        wide = gdf.pivot(index="date", columns="series_id", values="value").sort_index()
        wide.index = pd.DatetimeIndex(wide.index)
        origins = iter_origins(wide.index, split, horizon)
        if not origins:
            print(f"[{group}] no usable origins for split {split!r}, skipping")
            continue
        print(f"\n[{group}] {len(wide)} months, {wide.shape[1]} series, "
              f"{len(origins)} origins: {origins[0]:%Y-%m} .. {origins[-1]:%Y-%m}")

        # MASE scale, per origin, from the estimation sample only.
        for origin in origins:
            for sid in wide.columns:
                hist = wide.loc[:origin, sid].dropna().to_numpy(float)
                scales[(origin, sid)] = seasonal_mae_scale(hist, SEASON)

        # ---- univariate baselines and econometric models -------------------- #
        for name in [m for m in models if m in UNIVARIATE_MODELS]:
            t0 = time.time()
            print(f"  {name:<8}", end="", flush=True)
            warm_cache = {}  # (name, sid) -> previous origin's fitted parameters
            for i, origin in enumerate(origins):
                for sid in wide.columns:
                    hist = wide.loc[:origin, sid].dropna()
                    if len(hist) < SEASON * 3:
                        continue
                    try:
                        point, quant, warm = UNIVARIATE_MODELS[name](
                            hist, horizon, warm_cache.get(sid)
                        )
                    except Exception as exc:  # a single bad fit must not kill the run
                        print(f"\n    {name} failed at {origin:%Y-%m} {sid}: {exc}")
                        continue
                    if warm is not None:
                        warm_cache[sid] = warm
                    all_rows += collect_rows((origin, sid), name, point, quant,
                                             wide, horizon)
                if (i + 1) % 6 == 0:
                    print(".", end="", flush=True)
            print(f" {time.time() - t0:5.1f}s")

        # ---- VAR ------------------------------------------------------------ #
        if "var" in models and wide.shape[1] > 1:
            t0 = time.time()
            print(f"  {'var':<8}", end="", flush=True)
            for origin in origins:
                hist = wide.loc[:origin].dropna()
                if len(hist) < SEASON * 4:
                    continue
                try:
                    per_series = forecast_var(hist, horizon)
                except Exception as exc:
                    print(f"\n    var failed at {origin:%Y-%m}: {exc}")
                    continue
                for sid, (point, quant) in per_series.items():
                    all_rows += collect_rows((origin, sid), "var", point, quant,
                                             wide, horizon)
            print(f" {time.time() - t0:5.1f}s")

        # ---- TimesFM -------------------------------------------------------- #
        tfm_wanted = [m for m in models if m.startswith("timesfm")]
        if tfm_wanted:
            forecaster = get_forecaster(device, checkpoint, batch_size)
            if "timesfm_uni" in tfm_wanted:
                for key, point, quant in run_timesfm_univariate(
                    wide, origins, horizon, context_length, forecaster
                ):
                    all_rows += collect_rows(key, "timesfm_uni", point, quant,
                                             wide, horizon)
            if "timesfm_mv" in tfm_wanted and wide.shape[1] > 1:
                for key, point, quant in run_timesfm_multivariate(
                    wide, origins, horizon, context_length, forecaster
                ):
                    all_rows += collect_rows(key, "timesfm_mv", point, quant,
                                             wide, horizon)

    return pd.DataFrame(all_rows), scales


# --------------------------------------------------------------------------- #
# Reporting
# --------------------------------------------------------------------------- #
def plot_mase_by_horizon(by_h: pd.DataFrame, path: Path) -> None:
    agg = by_h.groupby(["model", "horizon"])["mase"].mean().reset_index()
    fig, ax = plt.subplots(figsize=(8, 5))
    for model, g in agg.groupby("model"):
        style = "-o" if model.startswith("timesfm") else "--"
        ax.plot(g.horizon, g.mase, style, label=model, linewidth=2 if style == "-o" else 1.4)
    ax.axhline(1.0, color="grey", linewidth=0.8)
    ax.annotate("seasonal naive, in-sample", (0.02, 1.02), xycoords=("axes fraction", "data"),
                fontsize=8, color="grey")
    ax.set_xlabel("forecast horizon (months)")
    ax.set_ylabel("MASE, averaged over series and origins")
    ax.set_title("Forecast accuracy by horizon")
    ax.legend(fontsize=8)
    fig.tight_layout()
    fig.savefig(path, dpi=140)
    plt.close(fig)


def plot_fanchart(forecasts: pd.DataFrame, panel: pd.DataFrame, sid: str,
                  path: Path, models=("timesfm_uni", "sarima")) -> None:
    g = forecasts[forecasts.series_id == sid]
    if g.empty:
        return
    origin = g.origin.max()
    hist = (panel[panel.series_id == sid].set_index("date")["value"]
            .sort_index().loc[:origin].tail(60))

    fig, ax = plt.subplots(figsize=(10, 5))
    ax.plot(hist.index, hist.values, color="black", linewidth=1.2, label="observed")

    truth = g[(g.origin == origin) & (g.model == g.model.iloc[0])]
    ax.plot(truth.target_date, truth.y_true, color="black", linewidth=1.2, linestyle=":",
            label="realised")

    colours = {"timesfm_uni": "tab:red", "timesfm_mv": "tab:orange", "sarima": "tab:blue",
               "ets": "tab:green", "rw": "tab:grey", "var": "tab:purple"}
    for model in models:
        m = g[(g.origin == origin) & (g.model == model)]
        if m.empty:
            continue
        colour = colours.get(model, None)
        ax.plot(m.target_date, m.y_pred, color=colour, linewidth=2, label=model)
        if {"q10", "q90"} <= set(m.columns) and m.q10.notna().any():
            ax.fill_between(m.target_date, m.q10, m.q90, color=colour, alpha=0.18,
                            label=f"{model} 80 %")

    ax.axvline(origin, color="grey", linestyle="--", linewidth=0.8)
    ax.set_title(f"{sid} — forecast from origin {origin:%Y-%m}")
    ax.set_ylabel("price index (2020 = 100)")
    ax.legend(fontsize=8)
    fig.tight_layout()
    fig.savefig(path, dpi=140)
    plt.close(fig)


def write_summary(by_h, overall, split, horizon, path: Path) -> None:
    ranked = (overall.groupby("model")[["mase", "rmse", "pinball", "coverage_80"]]
              .mean().sort_values("mase"))
    lines = [
        f"# Backtest summary — split `{split}`, horizon 1..{horizon}",
        "",
        "## Ranking, averaged over series, origins and horizons",
        "",
        "| model | MASE | RMSE | pinball | 80 % coverage |",
        "|---|---:|---:|---:|---:|",
    ]
    for model, r in ranked.iterrows():
        cov = "" if not np.isfinite(r.coverage_80) else f"{r.coverage_80:.2f}"
        pin = "" if not np.isfinite(r.pinball) else f"{r.pinball:.3f}"
        lines.append(f"| {model} | {r.mase:.3f} | {r.rmse:.2f} | {pin} | {cov} |")

    lines += ["", "MASE below 1 beats the in-sample seasonal naive. Coverage should",
              "sit near 0.80 — a model that is accurate but overconfident is not",
              "usable for decisions, and this column is where that shows up.", ""]

    h_max = int(by_h.horizon.max())
    lines += [f"## MASE at h=1 and h={h_max}", "",
              f"| model | h=1 | h={h_max} |", "|---|---:|---:|"]
    piv = by_h.groupby(["model", "horizon"])["mase"].mean().unstack()
    for model in ranked.index:
        a = piv.loc[model, 1] if 1 in piv.columns else np.nan
        b = piv.loc[model, h_max] if h_max in piv.columns else np.nan
        lines.append(f"| {model} | {a:.3f} | {b:.3f} |")

    dm_cols = [c for c in overall.columns if c.startswith("dm_pvalue")]
    if dm_cols:
        lines += ["", "## Diebold-Mariano against the random walk", "",
                  "Negative statistic: the model beats the random walk. "
                  "Origins overlap, so p-values are optimistic — read them as a "
                  "rough screen, not as proof.", "",
                  "| model | series | " + " | ".join(
                      c.replace("dm_pvalue_", "p ") for c in dm_cols) + " |",
                  "|---|---|" + "---:|" * len(dm_cols)]
        for _, r in overall[overall.model != "rw"].iterrows():
            ps = " | ".join(
                "" if not np.isfinite(r[c]) else f"{r[c]:.3f}" for c in dm_cols)
            lines.append(f"| {r.model} | {r.series_id} | {ps} |")

    if split == "stress":
        lines += ["", "## Regime-shift origins", "",
                  "| origin | event |", "|---|---|"]
        for d, label in STRESS_ORIGINS.items():
            lines.append(f"| {pd.Timestamp(d):%Y-%m} | {label} |")

    if split != "honest":
        lines += ["", "## Caveat", "",
                  "TimesFM 3.0 was released in August 2026 and its pretraining "
                  "corpus is undocumented at series level. These Destatis indices "
                  "are public, so origins before 2026 may be contaminated. Run "
                  "`--split honest` for the post-cutoff comparison and report both."]

    path.write_text("\n".join(lines) + "\n", encoding="utf-8")


# --------------------------------------------------------------------------- #
def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--panel", type=Path, default=PANEL)
    ap.add_argument("--models", default="rw,snaive,drift,ets,sarima,var",
                    help="comma separated, or 'all', or 'baselines'")
    ap.add_argument("--split", choices=["dev", "test", "stress", "honest"], default="test")
    ap.add_argument("--horizon", type=int, default=12)
    ap.add_argument("--context-length", type=int, default=512,
                    help="months of context handed to TimesFM (max 15360)")
    ap.add_argument("--device", default="cpu", help="cpu or cuda")
    ap.add_argument("--checkpoint", default="google/timesfm-3.0-pytorch")
    ap.add_argument("--batch-size", type=int, default=32)
    ap.add_argument("--group", choices=["forestry", "agriculture"], default=None)
    ap.add_argument("--out", type=Path, default=OUTPUT)
    args = ap.parse_args()

    if not args.panel.exists():
        sys.exit(f"No panel at {args.panel}. Run: python src/get_data.py --source synthetic")

    all_models = list(UNIVARIATE_MODELS) + ["var", "timesfm_uni", "timesfm_mv"]
    if args.models == "all":
        models = all_models
    elif args.models == "baselines":
        models = list(UNIVARIATE_MODELS) + ["var"]
    else:
        models = [m.strip() for m in args.models.split(",") if m.strip()]
    unknown = [m for m in models if m not in all_models]
    if unknown:
        sys.exit(f"Unknown model(s) {unknown}. Known: {all_models}")

    panel = pd.read_csv(args.panel, parse_dates=["date"])
    if args.group:
        panel = panel[panel.group == args.group]
    print(f"Panel: {panel.date.min():%Y-%m} .. {panel.date.max():%Y-%m}, "
          f"{panel.series_id.nunique()} series")
    print(f"Split: {args.split}, horizon 1..{args.horizon}, models: {models}")

    forecasts, scales = run_split(
        panel, models, args.split, args.horizon, args.context_length,
        args.device, args.checkpoint, args.batch_size,
    )
    if forecasts.empty:
        sys.exit("No forecasts were produced — check the panel length and the split.")

    args.out.mkdir(parents=True, exist_ok=True)
    suffix = f"_{args.split}"
    forecasts.to_csv(args.out / f"forecasts{suffix}.csv", index=False)

    by_h = metrics_by_horizon(forecasts, scales)
    overall = metrics_overall(forecasts, scales, reference_model="rw")
    by_h.to_csv(args.out / f"metrics_by_horizon{suffix}.csv", index=False)
    overall.to_csv(args.out / f"metrics_overall{suffix}.csv", index=False)

    plot_mase_by_horizon(by_h, args.out / f"mase_by_horizon{suffix}.png")
    plotted = [m for m in ("timesfm_uni", "timesfm_mv", "sarima", "rw") if m in models]
    for sid in forecasts.series_id.unique():
        plot_fanchart(forecasts, panel, sid,
                      args.out / f"fanchart_{sid}{suffix}.png", models=plotted)

    write_summary(by_h, overall, args.split, args.horizon,
                  args.out / f"summary{suffix}.md")

    print("\n" + "=" * 66)
    print(f"Ranking on split {args.split!r} (lower MASE is better)")
    print("=" * 66)
    ranked = (by_h.groupby("model")[["mase", "rmse", "pinball", "coverage_80"]]
              .mean().sort_values("mase"))
    print(ranked.to_string(float_format=lambda v: f"{v:8.3f}"))
    print(f"\nWrote forecasts, metrics, plots and summary{suffix}.md to {args.out}")


if __name__ == "__main__":
    main()
+321 −0
Changes for events/20260917/src/get_data.py: 321 added lines, 0 removed lines.
Original line number Diff line number Diff line
"""Build the monthly price-index panel for the TimesFM 3.0 experiment.

Three input paths, because GENESIS closed its anonymous account:

    --source synthetic   generated stand-in panel, no network, no account
    --source csv         a flat CSV downloaded by hand from GENESIS
    --source genesis     REST API, needs GENESIS_USERNAME / GENESIS_PASSWORD

All three write the same tidy long table to data/processed/prices_monthly.csv:

    date,series_id,group,value

Usage:
    python src/get_data.py --source synthetic
    python src/get_data.py --source csv --csv-file data/raw/61231.csv
    GENESIS_USERNAME=... GENESIS_PASSWORD=... python src/get_data.py --source genesis
"""

from __future__ import annotations

import argparse
import io
import os
import sys
from pathlib import Path

import numpy as np
import pandas as pd

HERE = Path(__file__).resolve().parent.parent
RAW = HERE / "data" / "raw"
PROCESSED = HERE / "data" / "processed"

GENESIS_BASE = "https://www-genesis.destatis.de/genesisWS/rest/2020"

# Destatis GENESIS tables named in the work item.
TABLES = {
    "61231-0002": "forestry",     # producer price index, forestry products
    "61211-0002": "agriculture",  # producer price index, agricultural products
}

# The series we want out of those tables, and the label fragments that identify
# them. GENESIS labels are verbose, so we match on a fragment, case-insensitively.
SERIES = {
    "forestry": {
        "oak_stemwood": ["oak", "eiche"],
        "beech_stemwood": ["beech", "buche"],
        "spruce_stemwood": ["spruce", "fichte"],
        "pine_stemwood": ["pine", "kiefer"],
    },
    "agriculture": {
        "wheat": ["wheat", "weizen"],
        "barley": ["barley", "gerste"],
        "sugar_beet": ["sugar beet", "zuckerrüben", "zuckerrueben"],
        "potatoes": ["potato", "kartoffel"],
    },
}


# --------------------------------------------------------------------------- #
# 1. Synthetic panel
# --------------------------------------------------------------------------- #
def make_synthetic(start="2005-01-01", end="2026-07-01", seed=20260917) -> pd.DataFrame:
    """A stand-in panel with the shape, seasonality and shocks of the real data.

    Not a simulation of German price formation. Its job is to let the backtest
    harness be written, debugged and demonstrated before anyone has downloaded
    anything, and to give the regime-shift stress set something to bite on.
    """
    rng = np.random.default_rng(seed)
    dates = pd.date_range(start, end, freq="MS")
    n = len(dates)
    t = np.arange(n)
    year = dates.year + (dates.month - 1) / 12.0

    def window(lo, hi):
        """Smooth 0->1->0 bump over the calendar window [lo, hi]."""
        mid, half = (lo + hi) / 2, (hi - lo) / 2
        return np.exp(-(((year - mid) / (half / 2.0)) ** 2))

    # Shocks shared across the panel, with different loadings per series.
    energy_spike = window(2021.5, 2023.0) * 45.0      # gas, diesel, fertiliser
    covid = window(2020.1, 2020.8) * -6.0
    beetle = window(2018.4, 2021.0) * -28.0           # bark-beetle calamity
    common = np.cumsum(rng.normal(0.0, 0.45, n))      # slow common drift

    spec = {
        # series_id: (group, level_2020, trend/yr, seas_amp, seas_phase,
        #             energy_load, beetle_load, sigma, ar1)
        "oak_stemwood":    ("forestry", 100, 1.6, 2.0, 0.1, 0.35, 0.20, 1.3, 0.55),
        "beech_stemwood":  ("forestry", 100, 1.2, 2.5, 0.2, 0.40, 0.45, 1.5, 0.55),
        "spruce_stemwood": ("forestry", 100, 0.6, 3.0, 0.0, 0.90, 1.00, 2.6, 0.60),
        "pine_stemwood":   ("forestry", 100, 0.8, 2.2, 0.0, 0.70, 0.75, 2.0, 0.58),
        "wheat":           ("agriculture", 100, 1.1, 4.0, 0.55, 0.80, 0.0, 2.4, 0.50),
        "barley":          ("agriculture", 100, 1.0, 4.5, 0.55, 0.75, 0.0, 2.6, 0.50),
        "sugar_beet":      ("agriculture", 100, 1.3, 1.5, 0.75, 0.45, 0.0, 1.8, 0.45),
        "potatoes":        ("agriculture", 100, 1.4, 12.0, 0.30, 0.35, 0.0, 5.5, 0.40),
    }

    frames = []
    for sid, (grp, lvl, slope, amp, phase, e_load, b_load, sigma, ar1) in spec.items():
        seasonal = amp * np.sin(2 * np.pi * ((dates.month - 1) / 12.0 + phase))
        # Second harmonic: real price seasonality is not a clean sine.
        seasonal += 0.4 * amp * np.sin(4 * np.pi * ((dates.month - 1) / 12.0 + phase))

        # AR(1) idiosyncratic noise, so first differences are not white.
        eps = rng.normal(0.0, sigma, n)
        noise = np.zeros(n)
        for i in range(1, n):
            noise[i] = ar1 * noise[i - 1] + eps[i]

        value = (
            lvl
            + slope * (year - 2020.0)
            + seasonal
            + e_load * energy_spike
            + b_load * beetle
            + covid
            + 0.6 * common
            + noise
        )
        frames.append(
            pd.DataFrame({"date": dates, "series_id": sid, "group": grp, "value": value})
        )

    out = pd.concat(frames, ignore_index=True)
    out["value"] = out["value"].clip(lower=5.0).round(1)
    return out


# --------------------------------------------------------------------------- #
# 2. GENESIS REST API
# --------------------------------------------------------------------------- #
def fetch_genesis(table: str, startyear: int = 2005) -> str:
    """Download one GENESIS table as a flat CSV (ffcsv) and return the text.

    Verified 2026-09-17: the anonymous GAST account no longer works, the API
    answers 401 Code 15 without credentials. Registration at
    https://www-genesis.destatis.de is free.
    """
    import requests

    user = os.environ.get("GENESIS_USERNAME", "")
    password = os.environ.get("GENESIS_PASSWORD", "")
    token = os.environ.get("GENESIS_TOKEN", "")
    if not token and not (user and password):
        sys.exit(
            "No GENESIS credentials. Set GENESIS_USERNAME and GENESIS_PASSWORD "
            "(or GENESIS_TOKEN), or use --source csv / --source synthetic."
        )

    headers = {"Authorization": f"Bearer {token}"} if token else {
        "username": user, "password": password
    }
    params = {
        "name": table,
        "area": "all",
        "compress": "false",
        "transpose": "false",
        "startyear": str(startyear),
        "format": "ffcsv",
        "language": "en",
    }
    resp = requests.post(
        f"{GENESIS_BASE}/data/tablefile", headers=headers, data=params, timeout=180
    )
    resp.raise_for_status()
    if resp.text.lstrip().startswith("{"):  # GENESIS reports errors as JSON 200s
        sys.exit(f"GENESIS error for {table}: {resp.text[:300]}")
    return resp.text


# --------------------------------------------------------------------------- #
# 3. Parse a GENESIS flat CSV
# --------------------------------------------------------------------------- #
def parse_genesis_csv(text: str, group: str) -> pd.DataFrame:
    """Parse a GENESIS ffcsv export into the tidy panel.

    GENESIS exports are semicolon separated with German decimals, and the exact
    column names depend on the table and the export options. This parser sniffs
    the time and value columns rather than hard-coding them, and prints what it
    found so a mismatch is visible immediately instead of silently producing an
    empty panel.
    """
    df = pd.read_csv(io.StringIO(text), sep=";", decimal=",", dtype=str)
    df.columns = [c.strip() for c in df.columns]
    print(f"  columns: {list(df.columns)}")

    time_col = next(
        (c for c in df.columns if c.lower() in ("time", "zeit", "time_label")), None
    )
    value_col = next(
        (c for c in df.columns if c.lower() in ("value", "wert", "value_variable_code")),
        None,
    )
    # The month sits in a "1_Auspraegung_Label"-style column, or in the time column.
    month_col = next(
        (c for c in df.columns if "auspraegung_label" in c.lower() or "month" in c.lower()),
        None,
    )
    label_cols = [c for c in df.columns if "label" in c.lower() and c != time_col]

    if time_col is None or value_col is None:
        sys.exit(
            "Could not identify the time and value columns in this export.\n"
            f"Columns seen: {list(df.columns)}\n"
            "Adjust parse_genesis_csv() — this is a five-minute fix once we see "
            "the real file."
        )

    month_map = {m: i + 1 for i, m in enumerate(
        ["january", "february", "march", "april", "may", "june", "july",
         "august", "september", "october", "november", "december"]
    )}
    month_map.update({m: i + 1 for i, m in enumerate(
        ["januar", "februar", "märz", "april", "mai", "juni", "juli",
         "august", "september", "oktober", "november", "dezember"]
    )})

    def to_date(row):
        year = str(row[time_col]).strip()[:4]
        raw = str(row[month_col]).strip().lower() if month_col else ""
        month = month_map.get(raw)
        if month is None:  # "MONAT01", "M01", "01"
            digits = "".join(ch for ch in raw if ch.isdigit())
            month = int(digits) if digits else None
        if month is None or not year.isdigit():
            return pd.NaT
        return pd.Timestamp(int(year), int(month), 1)

    df["date"] = df.apply(to_date, axis=1)
    df["value"] = pd.to_numeric(
        df[value_col].astype(str).str.replace(",", ".", regex=False), errors="coerce"
    )

    # Map the verbose GENESIS labels onto our short series ids.
    def to_series_id(row):
        haystack = " ".join(str(row[c]).lower() for c in label_cols)
        for sid, fragments in SERIES[group].items():
            if any(f in haystack for f in fragments):
                return sid
        return None

    df["series_id"] = df.apply(to_series_id, axis=1)
    df["group"] = group

    out = df.dropna(subset=["date", "value", "series_id"])
    out = out[["date", "series_id", "group", "value"]].drop_duplicates(
        subset=["date", "series_id"]
    )
    print(f"  kept {len(out)} rows, series: {sorted(out.series_id.unique())}")
    return out.reset_index(drop=True)


# --------------------------------------------------------------------------- #
def report(panel: pd.DataFrame) -> None:
    print("\nPanel summary")
    print("-" * 60)
    for (grp, sid), g in panel.groupby(["group", "series_id"]):
        print(
            f"  {grp:<12} {sid:<18} {g.date.min():%Y-%m} .. {g.date.max():%Y-%m}  "
            f"n={len(g):>4}  mean={g.value.mean():7.1f}"
        )
    print("-" * 60)
    print(f"  last observation T = {panel.date.max():%Y-%m}")
    print("  the backtest derives every split date from T, nothing is hard-coded")


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--source", choices=["synthetic", "csv", "genesis"],
                    default="synthetic")
    ap.add_argument("--csv-file", type=Path, nargs="*", default=[],
                    help="one or more GENESIS csv exports (--source csv)")
    ap.add_argument("--group", choices=["forestry", "agriculture"], default=None,
                    help="group of the --csv-file(s); inferred from the name if absent")
    ap.add_argument("--startyear", type=int, default=2005)
    ap.add_argument("--out", type=Path,
                    default=PROCESSED / "prices_monthly.csv")
    args = ap.parse_args()

    RAW.mkdir(parents=True, exist_ok=True)
    PROCESSED.mkdir(parents=True, exist_ok=True)

    if args.source == "synthetic":
        print("Generating the synthetic stand-in panel (no network, no account).")
        panel = make_synthetic()

    elif args.source == "genesis":
        frames = []
        for table, group in TABLES.items():
            print(f"GENESIS {table} ({group}) ...")
            text = fetch_genesis(table, args.startyear)
            (RAW / f"{table}.csv").write_text(text, encoding="utf-8")
            frames.append(parse_genesis_csv(text, group))
        panel = pd.concat(frames, ignore_index=True)

    else:  # csv
        if not args.csv_file:
            sys.exit("--source csv needs at least one --csv-file")
        frames = []
        for path in args.csv_file:
            group = args.group
            if group is None:
                group = "agriculture" if "61211" in path.name else "forestry"
            print(f"Parsing {path} as {group} ...")
            frames.append(
                parse_genesis_csv(path.read_text(encoding="utf-8", errors="replace"), group)
            )
        panel = pd.concat(frames, ignore_index=True)

    panel = panel.sort_values(["group", "series_id", "date"]).reset_index(drop=True)
    args.out.parent.mkdir(parents=True, exist_ok=True)
    panel.to_csv(args.out, index=False, date_format="%Y-%m-%d")
    report(panel)
    print(f"\nWrote {len(panel)} rows to {args.out}")


if __name__ == "__main__":
    main()
Loading
Loading