MacroCast / forecast.py
shubhranshu's picture
Fix src import on HF cache: use absolute() not resolve() (symlink)
bcb68fb verified
Raw
History Blame Contribute Delete
10.9 kB
#!/usr/bin/env python3
"""Zero-shot forecasting with MacroCast (or Forge) on arbitrary datasets.
MacroCast is the fine-tuned ~1.2M-parameter GatedDeltaProduct linear-RNN model. It forecasts in a single
forward pass (no windowing): give it a history of shape [T, N] (T timesteps, N
channels/variables) and it returns quantile forecasts for the next H steps.
Usage as a library
-------------------
from forecast import MacroCastForecaster
f = MacroCastForecaster(checkpoint="model.pth", config="config.yaml")
point, quantiles = f.predict(history_array, horizon=12) # history: [T, N]
Usage as a CLI
--------------
python inference/forecast.py \
--checkpoint /path/to/MacroCast/checkpoints/ckpt_2023-12.pth \
--config forge/config_forge_1M.yaml \
--csv my_series.csv \
--horizon 12 \
--output forecast.csv
The input CSV should have one row per timestep; the first column may be a date
index (auto-detected) and every remaining numeric column is treated as a channel.
"""
import sys
import json
import argparse
from pathlib import Path
import numpy as np
import pandas as pd
import torch
import yaml
# Make the vendored TempoPFN package importable as `src.*`.
# Works both in the CODE-SHARE layout (inference/forecast.py + ../tempopfn) and
# in a flat HF repo (forecast.py + ./tempopfn at the same level).
#
# NOTE: use .absolute() and NOT .resolve(): on the Hugging Face Hub cache, files in
# snapshots/<commit>/ are *symlinks* into blobs/. .resolve() would follow the symlink
# into blobs/ (where `tempopfn/` does not sit beside the file), breaking the import.
_HERE = Path(__file__).absolute().parent
for _cand in (_HERE, _HERE / "tempopfn", _HERE.parent, _HERE.parent / "tempopfn"):
if (_cand / "src").is_dir():
sys.path.insert(0, str(_cand))
break
if hasattr(torch, "compile"):
torch.compile = lambda fn=None, *a, **kw: (lambda inner: inner) if fn is None else fn
from src.data.containers import BatchTimeSeriesContainer # noqa: E402
from src.data.frequency import Frequency # noqa: E402
from src.models.model import TimeSeriesModel # noqa: E402
# Map common pandas-style freq strings to the model's Frequency enum.
# (The enum uses "A" for annual; there is no "Y".)
_FREQ_ALIASES = {
"M": "M", "MS": "M", "ME": "M",
"D": "D", "W": "W", "H": "H",
"Q": "Q", "QS": "Q",
"A": "A", "Y": "A", "YS": "A", "AS": "A",
}
def _to_frequency(freq: str):
name = _FREQ_ALIASES.get(freq.upper(), freq.upper())
return getattr(Frequency, name, Frequency.M)
def _fetch(source, filename):
"""Return a local path to `filename` from a local dir or an HF repo id.
For a local directory we just join the path. For a Hugging Face repo id we
download **only that single file** (cached) via `hf_hub_download` — never a
full-repo snapshot — so picking one vintage doesn't pull every checkpoint.
Raises FileNotFoundError if the file is absent.
"""
p = Path(source).expanduser()
if p.is_dir():
local = p / filename
if not local.exists():
raise FileNotFoundError(f"{filename} not found in {p}")
return local
from huggingface_hub import hf_hub_download
from huggingface_hub.utils import EntryNotFoundError, RepositoryNotFoundError
try:
return Path(hf_hub_download(repo_id=str(source), filename=filename))
except (EntryNotFoundError, RepositoryNotFoundError) as e:
raise FileNotFoundError(f"{filename} not found in repo {source!r}") from e
def _pick_vintage(manifest, year):
"""Pick the manifest vintage entry for `year` (exact, else nearest <=, else earliest)."""
vintages = [v for v in manifest.get("vintages", []) if v.get("year") is not None]
if not vintages:
return None
if year is None:
year = manifest.get("default_year")
for v in vintages:
if v.get("year") == year:
return v
le = [v for v in vintages if v["year"] <= (year if year is not None else 1 << 30)]
return max(le, key=lambda v: v["year"]) if le else min(vintages, key=lambda v: v["year"])
class MacroCastForecaster:
"""Load a MacroCast/Forge checkpoint and produce zero-shot forecasts."""
@staticmethod
def available_years(source):
"""List the vintage years available in a local dir or HF repo.
Only `manifest.json` is read/downloaded — not the checkpoints.
"""
try:
manifest = _fetch(source, "manifest.json")
except FileNotFoundError:
return []
m = json.loads(Path(manifest).read_text())
return [v["year"] for v in m.get("vintages", []) if v.get("year") is not None]
@classmethod
def from_pretrained(cls, source, year=None, config=None, device=None):
"""Load MacroCast for a chosen vintage `year` from a local dir or HF repo id.
Pass the Hugging Face repo id directly, e.g.::
model = MacroCastForecaster.from_pretrained("user/MacroCast", year=2023)
Reads `manifest.json` to map year -> checkpoint and downloads **only** the
matching vintage file (plus `config.yaml`) — not the whole repo. If `year`
is None, the latest (`default_year`) vintage is used. Falls back to a single
`model.pth` if the repo/dir has no manifest.
"""
try:
manifest_path = _fetch(source, "manifest.json")
except FileNotFoundError:
manifest_path = None
if manifest_path is not None:
m = json.loads(Path(manifest_path).read_text())
chosen = _pick_vintage(m, year)
if chosen is None:
raise FileNotFoundError(f"No vintage models found for source={source!r}")
cfg = Path(config) if config else _fetch(source, m.get("config", "config.yaml"))
ckpt = _fetch(source, chosen["file"])
print(f"Loading MacroCast vintage year={chosen.get('year')} "
f"(cutoff={chosen.get('cutoff')}, n_vars={chosen.get('n_vars')})")
return cls(checkpoint=ckpt, config=cfg, device=device)
# No manifest → single-model repo/dir.
cfg = Path(config) if config else _fetch(source, "config.yaml")
single = _fetch(source, "model.pth")
return cls(checkpoint=single, config=cfg, device=device)
def __init__(self, checkpoint, config, device=None):
self.device = torch.device(
device or ("cuda" if torch.cuda.is_available() else "cpu"))
with open(config) as f:
cfg = yaml.safe_load(f)["TimeSeriesModel"]
self.model = TimeSeriesModel(**cfg).to(self.device)
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
state = ckpt["model_state_dict"] if isinstance(ckpt, dict) and \
"model_state_dict" in ckpt else ckpt
# MacroCast checkpoints have no adapter weights; load non-strict to be safe.
self.model.load_state_dict(state, strict=False)
self.model.eval()
# Variable names the model was fine-tuned on (informational).
self.cols = ckpt.get("cols") if isinstance(ckpt, dict) else None
self.quantiles = list(self.model.quantiles)
@torch.no_grad()
def predict(self, history, horizon=12, freq="M", start="1970-01"):
"""Forecast `horizon` steps ahead.
Parameters
----------
history : np.ndarray, shape [T] or [T, N]
Observed history. N channels are forecast jointly.
horizon : int
freq : str, one of M, D, W, H, Q, Y
start : str, start timestamp of the history (for time features)
Returns
-------
point : np.ndarray [horizon, N] median (0.5-quantile) forecast
quantiles : np.ndarray [horizon, N, Q] all quantile forecasts
"""
hist = np.asarray(history, dtype=np.float32)
if hist.ndim == 1:
hist = hist[:, None]
T, N = hist.shape
fut = np.zeros((horizon, N), dtype=np.float32)
fr = _to_frequency(freq)
batch = BatchTimeSeriesContainer(
history_values=torch.tensor(hist[None], dtype=torch.float32, device=self.device),
future_values=torch.tensor(fut[None], dtype=torch.float32, device=self.device),
start=[np.datetime64(start)],
frequency=[fr],
)
autocast = self.device.type == "cuda"
with torch.autocast(device_type=self.device.type, dtype=torch.bfloat16,
enabled=autocast):
out = self.model(batch)
pred = self.model.scaler.inverse_scale(out["result"].float(),
out["scale_statistics"])
arr = pred.cpu().numpy()[0] # [horizon, N, Q]
q_idx = self.quantiles.index(0.5) if 0.5 in self.quantiles else arr.shape[-1] // 2
point = arr[:, :, q_idx] # [horizon, N]
return point.astype(np.float32), arr.astype(np.float32)
def _read_csv(path):
df = pd.read_csv(path)
# Auto-detect a date column as the index.
first = df.columns[0]
try:
idx = pd.to_datetime(df[first])
df = df.drop(columns=[first]).set_index(idx)
except (ValueError, TypeError):
pass
df = df.select_dtypes(include=[np.number])
return df
def main():
ap = argparse.ArgumentParser(description="Zero-shot forecasting with MacroCast.")
ap.add_argument("--checkpoint", required=True, help="MacroCast/Forge .pth checkpoint")
ap.add_argument("--config", required=True, help="Architecture YAML (forge/config_forge_1M.yaml)")
ap.add_argument("--csv", required=True, help="Input CSV: rows=time, cols=variables")
ap.add_argument("--horizon", type=int, default=12)
ap.add_argument("--freq", default="M")
ap.add_argument("--output", default="forecast.csv")
ap.add_argument("--device", default=None)
args = ap.parse_args()
df = _read_csv(args.csv)
print(f"Loaded {df.shape[0]} timesteps x {df.shape[1]} channels from {args.csv}")
f = MacroCastForecaster(args.checkpoint, args.config, device=args.device)
start = str(df.index[0].to_datetime64()) if isinstance(df.index, pd.DatetimeIndex) \
else "1970-01"
point, _ = f.predict(df.to_numpy(np.float32), horizon=args.horizon,
freq=args.freq, start=start)
# Build a forward index if the input was dated.
if isinstance(df.index, pd.DatetimeIndex):
future_idx = pd.date_range(df.index[-1], periods=args.horizon + 1,
freq=args.freq)[1:]
else:
future_idx = range(args.horizon)
out = pd.DataFrame(point, columns=df.columns, index=future_idx)
out.to_csv(args.output)
print(f"Wrote {args.horizon}-step forecast for {df.shape[1]} channels → {args.output}")
if __name__ == "__main__":
main()