trichronos-0.1B / evaluate.py
iravikr's picture
v0.1.0 — TriChronos-0.1B initial release
d3b9aa2 verified
Raw
History Blame Contribute Delete
8.7 kB
"""
evaluate.py — TriChronos-0.1B
Zero-shot MASE evaluation against the Monash Time Series Forecasting benchmark.
Designed to run on a free CPU tier — no GPU required.
Metrics
-------
MASE = MAE(forecast, actual) / naïve_seasonal_MAE
where naïve_seasonal_MAE uses the last known value from the context
as the forecast (seasonal period = 1, i.e. random-walk naive baseline).
For probabilistic forecasts we evaluate the **median quantile** (τ=0.50).
Usage
-----
python evaluate.py --checkpoint checkpoints/model_state.pt
"""
from __future__ import annotations
import argparse
from pathlib import Path
from typing import Dict, List, Optional, Tuple
import numpy as np
import torch
from torch.utils.data import DataLoader
from data_pipeline import (
_silver_normalize,
_gold_patch,
PATCH_SIZE,
FORECAST_HORIZON,
)
from model import TriChronos, QUANTILE_LEVELS
# ---------------------------------------------------------------------------
# Monash datasets to evaluate on
# ---------------------------------------------------------------------------
# We use the Parquet-based `autogluon/chronos_datasets` mirror rather than the
# original `monash_tsf`, which is a loading-*script* dataset and no longer
# loads (recent `datasets` versions dropped `trust_remote_code`). The mirror
# has one "train" split per config, each row = one univariate series in
# `target`.
MONASH_REPO: str = "autogluon/chronos_datasets"
# Each entry: (hf_dataset_name, config_name, seasonal_period)
MONASH_DATASETS: List[Tuple[str, Optional[str], int]] = [
(MONASH_REPO, "monash_m1_monthly", 12),
(MONASH_REPO, "monash_m1_quarterly", 4),
(MONASH_REPO, "monash_m1_yearly", 1),
(MONASH_REPO, "monash_m3_monthly", 12),
(MONASH_REPO, "monash_m3_quarterly", 4),
(MONASH_REPO, "monash_m3_yearly", 1),
(MONASH_REPO, "m4_monthly", 12),
(MONASH_REPO, "m4_quarterly", 4),
(MONASH_REPO, "m4_yearly", 1),
(MONASH_REPO, "monash_tourism_monthly", 12),
(MONASH_REPO, "monash_tourism_quarterly", 4),
(MONASH_REPO, "monash_tourism_yearly", 1),
(MONASH_REPO, "monash_electricity_hourly", 24),
(MONASH_REPO, "monash_traffic", 24),
(MONASH_REPO, "monash_weather", 1),
]
# Median is quantile index for τ=0.50
MEDIAN_IDX: int = QUANTILE_LEVELS.index(0.5)
# ---------------------------------------------------------------------------
# Naive seasonal baseline
# ---------------------------------------------------------------------------
def naive_mae(history: np.ndarray, horizon: int, period: int) -> float:
"""
MAE of the seasonal-naive forecast on the history (in-sample).
For period=1 (random walk): forecast[t] = history[t-1].
"""
if len(history) <= period:
return float(np.abs(history).mean()) or 1.0
errors = np.abs(history[period:] - history[:-period])
return float(errors.mean()) or 1.0
# ---------------------------------------------------------------------------
# Single-dataset evaluation
# ---------------------------------------------------------------------------
def evaluate_dataset(
model: TriChronos,
ds_name: str,
subset: Optional[str],
period: int,
device: torch.device,
max_series: int = 200,
) -> Tuple[float, int]:
"""
Run zero-shot evaluation on one Monash dataset.
Returns
-------
mean_mase : float — mean MASE across all evaluated series
n_series : int — number of series evaluated
"""
from datasets import load_dataset
try:
ds_kwargs = dict(
path=ds_name,
split="train", # chronos_datasets exposes only a "train" split (full series)
streaming=True, # stream so 48k-series M4 configs don't download in full
)
if subset:
ds_kwargs["name"] = subset
dataset = load_dataset(**ds_kwargs)
except Exception as exc:
print(f" [SKIP] {subset or ds_name}: {exc}")
return float("nan"), 0
mase_scores: List[float] = []
model.eval()
with torch.no_grad():
for i, example in enumerate(dataset):
if i >= max_series:
break
values = example.get("target", None)
if values is None or len(values) < PATCH_SIZE * 2 + FORECAST_HORIZON:
continue
values = np.asarray(values, dtype=np.float32)
# Normalise
normed, mean_, std_ = _silver_normalize(values)
# Use the last window as context; forecast FORECAST_HORIZON steps
# (if the series is shorter than FORECAST_HORIZON, skip it)
if len(normed) < FORECAST_HORIZON * 2:
continue
# Context = all but last FORECAST_HORIZON steps
context = normed[: -FORECAST_HORIZON]
actual_normed = normed[-FORECAST_HORIZON:]
# Build patches from context
n_full = len(context) // PATCH_SIZE
if n_full == 0:
continue
aligned = context[-n_full * PATCH_SIZE:]
patches = aligned.reshape(n_full, PATCH_SIZE)
# Pad/truncate to max_patches=64
max_patches = 64
if patches.shape[0] > max_patches:
patches = patches[-max_patches:]
elif patches.shape[0] < max_patches:
pad = np.zeros((max_patches - patches.shape[0], PATCH_SIZE), dtype=np.float32)
patches = np.concatenate([pad, patches], axis=0)
patches_t = torch.from_numpy(patches).unsqueeze(0).to(device) # (1, 64, 8)
# Model forward → (1, horizon, n_quantiles)
preds = model(patches_t)
# Median prediction
median_normed = preds[0, :FORECAST_HORIZON, MEDIAN_IDX].cpu().numpy()
# Compute MAE in normalised space
mae = float(np.abs(median_normed - actual_normed).mean())
# Naive baseline MAE (on normalised context)
n_mae = naive_mae(context, FORECAST_HORIZON, period)
mase = mae / max(n_mae, 1e-8)
mase_scores.append(mase)
if not mase_scores:
return float("nan"), 0
return float(np.mean(mase_scores)), len(mase_scores)
# ---------------------------------------------------------------------------
# Main evaluation loop
# ---------------------------------------------------------------------------
def evaluate(args: argparse.Namespace):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Device: {device}\n", flush=True)
# Load model
model = TriChronos(patch_size=PATCH_SIZE, horizon=FORECAST_HORIZON).to(device)
if args.checkpoint and Path(args.checkpoint).exists():
state = torch.load(args.checkpoint, map_location=device)
model.load_state_dict(state)
print(f"Loaded checkpoint: {args.checkpoint}\n", flush=True)
else:
print("⚠️ No checkpoint loaded — evaluating with random weights (sanity check only).\n")
model.eval()
# Per-dataset results
results: Dict[str, Dict] = {}
all_mase: List[float] = []
header = f"{'Dataset':<35} {'MASE':>8} {'N':>6}"
print(header)
print("-" * len(header))
for ds_name, subset, period in MONASH_DATASETS:
label = subset or ds_name
mase, n = evaluate_dataset(model, ds_name, subset, period, device, max_series=args.max_series)
results[label] = {"mase": mase, "n": n}
if not np.isnan(mase):
all_mase.append(mase)
print(f"{label:<35} {mase:>8.4f} {n:>6}")
else:
print(f"{label:<35} {'N/A':>8} {n:>6}")
print("-" * len(header))
if all_mase:
agg_mase = float(np.mean(all_mase))
print(f"{'AGGREGATE (mean MASE)':<35} {agg_mase:>8.4f}")
else:
print("No datasets were successfully evaluated.")
return results
# ---------------------------------------------------------------------------
# Entry point
# ---------------------------------------------------------------------------
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Evaluate TriChronos-0.1B on Monash")
p.add_argument(
"--checkpoint",
type=str,
default="checkpoints/model_state.pt",
help="Path to model checkpoint",
)
p.add_argument(
"--max-series",
type=int,
default=200,
help="Max series per dataset (limits CPU eval time)",
)
return p.parse_args()
if __name__ == "__main__":
evaluate(parse_args())