LiveHouse-TS / scripts /run_demo_eval.py
ziyuzhou02's picture
Deploy GitHub main 3feb6cda1511
e317359 verified
Raw History Blame Contribute Delete
5 kB
#!/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()