Spaces:
Running
Running
Download scripts/run_demo_eval.py from ThinkcatLab/LiveHouse-TS: direct link, hf CLI and curl.
- Browser
- Download file 5 kB
-
https://huggingface.co/spaces/ThinkcatLab/LiveHouse-TS/resolve/main/scripts/run_demo_eval.py
- Command line
-
hf download hf://spaces/ThinkcatLab/LiveHouse-TS/scripts/run_demo_eval.py
-
curl -L -o run_demo_eval.py https://huggingface.co/spaces/ThinkcatLab/LiveHouse-TS/resolve/main/scripts/run_demo_eval.py
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() | |