Add forecast galleries: per-model fan charts on held-out series + raw forecast dumps + renderer
4dbfe21 verified | #!/usr/bin/env python3 | |
| """Forecast fan-chart gallery, one facet per model. | |
| Rows = assets, cols = models; per timeframe one PNG. Each facet: recent | |
| real series (grey), the actual future (black), the model's median forecast | |
| (dashed) and its quantile bands (10-90% and 30-70% shaded; for a 13-level | |
| model the trained 2.5-97.5% tail band is shaded very lightly underneath). | |
| Quantile indices are taken from each dump's own `levels` array. | |
| Usage: zoo_plot2.py <tf> (reads data/forecast-dumps/plotdump-<tag>-<tf>.npz) | |
| """ | |
| import sys | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| TAGS = [("yv02", "Yumoto v0.2", "#0072B2"), | |
| ("yv01", "Yumoto v0.1", "#56B4E9"), | |
| ("toto25b", "Toto-2.0-2.5B", "#E69F00"), | |
| ("tirex2", "TiRex-2", "#009E73"), | |
| ("timesfm3", "TimesFM-3.0", "#CC79A7")] | |
| TF_TITLE = {"5m": "5-minute", "1h": "hourly", "4h": "4-hourly", "1d": "daily"} | |
| tf = sys.argv[1] | |
| dumps = {} | |
| for tag, _, _ in TAGS: | |
| try: | |
| dumps[tag] = np.load(f"data/forecast-dumps/plotdump-{tag}-{tf}.npz", | |
| allow_pickle=True) | |
| except FileNotFoundError: | |
| pass | |
| ref = next(iter(dumps.values())) | |
| assets = [str(x) for x in ref["assets"]] | |
| def q_at(d, name, level): | |
| lv = np.asarray(d["levels"], float) | |
| i = int(np.argmin(np.abs(lv - level))) | |
| assert abs(lv[i] - level) < 1e-6, (level, lv) | |
| return np.asarray(d[f"q_{name}"])[:, i] | |
| fig, axes = plt.subplots(len(assets), len(TAGS), | |
| figsize=(3.6 * len(TAGS), 2.7 * len(assets)), | |
| squeeze=False, dpi=150) | |
| for r, name in enumerate(assets): | |
| for c, (tag, label, color) in enumerate(TAGS): | |
| ax = axes[r][c] | |
| if tag not in dumps: | |
| ax.axis("off"); continue | |
| d = dumps[tag] | |
| ctx = np.asarray(d[f"ctx_{name}"], float) | |
| act = np.asarray(d[f"act_{name}"], float) | |
| H = act.size | |
| # show a context tail ~1.5x the horizon so the forecast has room | |
| keep = min(ctx.size, int(1.5 * H)) | |
| xs_c = np.arange(-keep, 0); xs_f = np.arange(H) | |
| lv = np.asarray(d["levels"], float) | |
| if lv.min() < 0.05: # 13-level model: show the trained tail band | |
| ax.fill_between(xs_f, q_at(d, name, 0.025), q_at(d, name, 0.975), | |
| color=color, alpha=0.10, lw=0) | |
| ax.fill_between(xs_f, q_at(d, name, 0.1), q_at(d, name, 0.9), | |
| color=color, alpha=0.22, lw=0) | |
| ax.fill_between(xs_f, q_at(d, name, 0.3), q_at(d, name, 0.7), | |
| color=color, alpha=0.38, lw=0) | |
| ax.plot(xs_c, ctx[-keep:], color="#999999", lw=0.9) | |
| ax.plot(xs_f, act, color="#111111", lw=1.1) | |
| ax.plot(xs_f, q_at(d, name, 0.5), color=color, lw=1.2, ls="--") | |
| ax.axvline(0, color="#bbbbbb", lw=0.7) | |
| lo = min(ctx[-keep:].min(), act.min()); hi = max(ctx[-keep:].max(), act.max()) | |
| pad = 0.35 * (hi - lo + 1e-12) | |
| ax.set_ylim(lo - pad, hi + pad) | |
| ax.tick_params(labelsize=6.5) | |
| for s in ("top", "right"): ax.spines[s].set_visible(False) | |
| if r == 0: ax.set_title(label, fontsize=10, color=color) | |
| if c == 0: ax.set_ylabel(name, fontsize=9) | |
| if r == len(assets) - 1: ax.set_xlabel("steps from forecast origin", fontsize=7.5) | |
| fig.suptitle(f"{TF_TITLE[tf]} candles — held-out forecasts " | |
| "(grey: context · black: what actually happened · " | |
| "dashed: median forecast · shading: quantile bands)", | |
| fontsize=11, y=1.005) | |
| fig.tight_layout() | |
| fig.savefig(f"plots/forecasts/{tf}.png", bbox_inches="tight") | |
| print("wrote", f"plots/forecasts/{tf}.png") | |