Spaces:
Running
Running
File size: 4,999 Bytes
e317359 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 | #!/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()
|