SPB Git forge

spb/wp12_uqo

Public
5commits 1branches 0releases
1.2 MBsize
maindefault branch
1 mo agolast push
Python 75.9% TeX 24%
6.0 KB · 169 lines python
Raw Blame History
1#!/usr/bin/env python32"""3================================================================4Auteur  : Simon-Pierre Boucher5Contact : contact@spboucher.ai6Projet  : Prévision de volatilité réalisée multi-actifs7          (HAR-RV vs GARCH vs Machine Learning)8Fichier : 03_run_models.py9Description : Étape 03 — Prévisions out-of-sample roulantes pour10              toutes les familles de modèles (HAR, GARCH, ML,11              deep, pooled), parallélisées par actif.12Usage : python scripts/03_run_models.py --family econ|ml|deep|pooled|all13        [--tickers SPY,BTC] [--workers 8]14================================================================15"""1617from __future__ import annotations1819import argparse20import logging21import sys22from concurrent.futures import ProcessPoolExecutor, as_completed23from pathlib import Path2425import pandas as pd2627sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))2829from wp12 import config  # noqa: E4023031logger = logging.getLogger("03_run_models")3233FC_DIR = config.path("processed") / "forecasts"3435HORIZONS = tuple(config.load_config()["evaluation"]["horizons"])36WINDOW = int(config.load_config()["evaluation"]["estimation_window_days"])37REFIT = int(config.load_config()["evaluation"]["reestimation_freq_days"])383940def _load_rv(ticker: str) -> pd.DataFrame:41    return pd.read_parquet(config.path("processed") / f"rv_{ticker}.parquet")424344def run_econ(ticker: str) -> str:45    """All HAR-family and GARCH-family forecasts for one ticker."""46    from wp12 import models_econ as me4748    df = _load_rv(ticker)49    parts = []50    for spec in ["HAR", "HAR-J", "HAR-CJ", "SHAR", "HARQ", "LogHAR"]:51        parts.append(me.rolling_har_forecasts(52            df, spec, horizons=HORIZONS, window=WINDOW, refit_every=REFIT))53    for kind in ["GARCH", "GJR", "EGARCH"]:54        parts.append(me.rolling_garch_forecasts(55            df, kind, horizons=HORIZONS, window=WINDOW, refit_every=REFIT))56    parts.append(me.rolling_realgarch_forecasts(57        df, horizons=HORIZONS, window=WINDOW, refit_every=REFIT))58    out = pd.concat(parts, ignore_index=True)59    out["ticker"] = ticker60    f = FC_DIR / f"econ_{ticker}.parquet"61    out.to_parquet(f, index=False)62    return f"{ticker}: {len(out)} econ forecasts"636465def run_ml(ticker: str) -> str:66    """Tabular ML forecasts (2 feature sets x 6 estimators) for one ticker."""67    from wp12 import models_ml as ml6869    panel = pd.read_parquet(config.path("processed") / "panel.parquet")70    frames = ml.ml_features(panel)71    frame = frames[ticker]72    parts = []73    for name in ["LASSO", "Ridge", "ElasticNet", "RF", "XGBoost", "LightGBM"]:74        for fs in ["har", "extended"]:75            parts.append(ml.rolling_ml_forecasts(76                frame, name, feature_set=fs, horizons=HORIZONS,77                window=WINDOW, refit_every=REFIT))78    out = pd.concat(parts, ignore_index=True)79    out["ticker"] = ticker80    f = FC_DIR / f"ml_{ticker}.parquet"81    out.to_parquet(f, index=False)82    return f"{ticker}: {len(out)} ML forecasts"838485def run_deep(ticker: str) -> str:86    """LSTM and Transformer forecasts for one ticker."""87    from wp12 import models_ml as ml8889    panel = pd.read_parquet(config.path("processed") / "panel.parquet")90    frames = ml.ml_features(panel)91    frame = frames[ticker]92    parts = [93        ml.rolling_deep_forecasts(frame, arch, horizons=HORIZONS, window=WINDOW)94        for arch in ["LSTM", "Transformer"]95    ]96    out = pd.concat(parts, ignore_index=True)97    out["ticker"] = ticker98    f = FC_DIR / f"deep_{ticker}.parquet"99    out.to_parquet(f, index=False)100    return f"{ticker}: {len(out)} deep forecasts"101102103def run_interpret() -> str:104    """SHAP values and permutation importance of the pooled LightGBM."""105    from wp12 import models_ml as ml106107    panel = pd.read_parquet(config.path("processed") / "panel.parquet")108    frames = ml.ml_features(panel)109    shap_tab, perm_tab = ml.shap_and_permutation(frames, h=1, window=WINDOW)110    out = config.path("reproduced")111    shap_tab.to_csv(out / "shap_h1.csv", index=False)112    perm_tab.to_csv(out / "permutation_h1.csv", index=False)113    return f"interpret: {len(shap_tab)} features"114115116def run_pooled() -> str:117    """Pooled multi-asset LightGBM (single job, all tickers jointly)."""118    from wp12 import models_ml as ml119120    panel = pd.read_parquet(config.path("processed") / "panel.parquet")121    frames = ml.ml_features(panel)122    out = ml.rolling_pooled_lgbm(123        frames, horizons=HORIZONS, window=WINDOW, refit_every=REFIT)124    f = FC_DIR / "pooled.parquet"125    out.to_parquet(f, index=False)126    return f"pooled: {len(out)} forecasts"127128129def main() -> None:130    """Dispatch model runs across tickers with a process pool."""131    ap = argparse.ArgumentParser()132    ap.add_argument("--family", default="all",133                    choices=["econ", "ml", "deep", "pooled", "interpret", "all"])134    ap.add_argument("--tickers", default=None, help="comma-separated subset")135    ap.add_argument("--workers", type=int, default=8)136    args = ap.parse_args()137138    config.setup_logging()139    FC_DIR.mkdir(parents=True, exist_ok=True)140    tickers = (141        args.tickers.split(",") if args.tickers142        else [i["ticker"] for i in config.universe()]143    )144    jobs: list[tuple] = []145    if args.family in ("econ", "all"):146        jobs += [(run_econ, tk) for tk in tickers]147    if args.family in ("ml", "all"):148        jobs += [(run_ml, tk) for tk in tickers]149    if args.family in ("deep", "all"):150        jobs += [(run_deep, tk) for tk in tickers]151152    if jobs:153        with ProcessPoolExecutor(max_workers=args.workers) as pool:154            futs = {pool.submit(fn, tk): (fn.__name__, tk) for fn, tk in jobs}155            for fut in as_completed(futs):156                name, tk = futs[fut]157                try:158                    logger.info("done %s %s -> %s", name, tk, fut.result())159                except Exception as exc:  # noqa: BLE001160                    logger.error("FAILED %s %s: %s", name, tk, exc)161    if args.family in ("pooled", "all"):162        logger.info(run_pooled())163    if args.family in ("interpret", "all"):164        logger.info(run_interpret())165166167if __name__ == "__main__":168    main()169