#!/usr/bin/env python3 """Run a demo evaluation on synthetic multivariate data and write leaderboard artifacts.""" from __future__ import annotations import argparse import csv import json from pathlib import Path from dotenv import load_dotenv from gluonts.model import evaluate_model from gluonts.time_feature import get_seasonality from tsfm_bench.data.dataset import Dataset from tsfm_bench.data.registry import load_data_source, load_dataset_properties from tsfm_bench.eval.metrics import RESULT_COLUMNS, get_eval_metrics from tsfm_bench.eval.predictors import SeasonalNaivePredictor def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( "--config", type=Path, default=Path("configs/datasets/synthetic_demo.yaml"), help="Dataset registry config path", ) parser.add_argument( "--model-name", default="TSFM2", help="Model name written to all_results.csv", ) parser.add_argument( "--output-dir", type=Path, default=Path("results/tsfm2"), help="Directory for all_results.csv and config.json", ) parser.add_argument( "--space-results-dir", type=Path, default=Path("space/results/tsfm2"), help="Mirror results into HF Space folder", ) return parser.parse_args() def build_config_name(ds_name: str, ds_key: str, frequency: str, term: str) -> str: return f"{ds_key}/{frequency}/{term}" def main() -> None: load_dotenv() args = parse_args() source = load_data_source(args.config) properties = load_dataset_properties(args.config) metrics = get_eval_metrics() args.output_dir.mkdir(parents=True, exist_ok=True) csv_path = args.output_dir / "all_results.csv" with csv_path.open("w", newline="") as csvfile: writer = csv.writer(csvfile) writer.writerow(RESULT_COLUMNS) for ds_name in source.list_datasets(): meta = source.get_metadata(ds_name) ds_key = ds_name.split("/")[0].lower() for term in meta.terms: to_univariate = meta.num_variates > 1 dataset = Dataset( name=ds_name, term=term, to_univariate=to_univariate, source=source, ) season_length = get_seasonality(dataset.freq) predictor = SeasonalNaivePredictor( prediction_length=dataset.prediction_length, season_length=season_length, quantile_levels=[0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9], ) res = evaluate_model( predictor, test_data=dataset.test_data, metrics=metrics, batch_size=64, axis=None, mask_invalid_label=True, allow_nan_forecast=False, seasonality=season_length, ) metric_value = lambda key: float(res[key].iloc[0]) writer.writerow( [ build_config_name(ds_name, ds_key, meta.frequency, term), args.model_name, metric_value("MSE[mean]"), metric_value("MSE[0.5]"), metric_value("MAE[0.5]"), metric_value("MASE[0.5]"), metric_value("MAPE[0.5]"), metric_value("sMAPE[0.5]"), metric_value("MSIS"), metric_value("RMSE[mean]"), metric_value("NRMSE[mean]"), metric_value("ND[0.5]"), metric_value("mean_weighted_sum_quantile_loss"), properties[ds_key]["domain"], properties[ds_key]["num_variates"], ] ) print(f"Evaluated {ds_name} ({term})") config = { "model": args.model_name, "model_type": "statistical", "model_dtype": "float32", "model_link": "https://github.com/zhouziyu02/TS-Live", "code_link": "https://github.com/zhouziyu02/TS-Live/blob/main/scripts/run_demo_eval.py", "org": "LiveHouse-TS", "testdata_leakage": "No", "replication_code_available": "Yes", } config_path = args.output_dir / "config.json" config_path.write_text(json.dumps(config, indent=4) + "\n") if args.space_results_dir != args.output_dir: args.space_results_dir.mkdir(parents=True, exist_ok=True) (args.space_results_dir / "all_results.csv").write_text(csv_path.read_text()) (args.space_results_dir / "config.json").write_text(config_path.read_text()) print(f"Wrote {csv_path}") if __name__ == "__main__": main()