Spaces:
Running
Running
| #!/usr/bin/env python3 | |
| """Run live TS-Bench + TSFM.ai evaluation inside the HF Space (or locally).""" | |
| from __future__ import annotations | |
| import csv | |
| import json | |
| import logging | |
| import os | |
| import sys | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from typing import Any | |
| import yaml | |
| SPACE_ROOT = Path(__file__).resolve().parents[1] | |
| if str(SPACE_ROOT) not in sys.path: | |
| sys.path.insert(0, str(SPACE_ROOT)) | |
| os.environ.setdefault("TS_BENCH_ROOT", str(SPACE_ROOT / "vendor" / "ts_bench")) | |
| from tsfm_bench.data.registry import load_data_source, load_dataset_properties | |
| from tsfm_bench.data.ts_bench import TsBenchDataSource | |
| from tsfm_bench.eval.api_predictor import DEFAULT_QUANTILES, TsfmApiConfig, TsfmApiPredictor | |
| from tsfm_bench.eval.online_eval import run_online_eval_for_dataset | |
| from src.eval_schedule import with_next_eval_fields | |
| from src.benchmark_config import BENCHMARK_INTERVAL_SECONDS | |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") | |
| logger = logging.getLogger(__name__) | |
| DATA_CONFIG = SPACE_ROOT / "configs" / "datasets" / "ts_bench.yaml" | |
| MODEL_CONFIG = SPACE_ROOT / "configs" / "models" / "online_tsfm.yaml" | |
| OUTPUT_ROOT = Path(os.getenv("TSFM_RESULTS_PATH", str(SPACE_ROOT / "results"))) | |
| EVAL_STATE_PATH = OUTPUT_ROOT / "eval_state.json" | |
| def model_display_name(model_spec: dict[str, Any]) -> str: | |
| return model_spec.get("display_name") or model_spec["model_id"] | |
| def model_output_slug(model_spec: dict[str, Any]) -> str: | |
| import re | |
| name = model_spec.get("display_name") or model_spec["model_id"] | |
| return re.sub(r"[^a-zA-Z0-9]+", "_", name).strip("_").lower() | |
| def load_model_specs(path: Path) -> list[dict[str, Any]]: | |
| return yaml.safe_load(path.read_text()).get("models", []) | |
| def write_model_results( | |
| model_spec: dict[str, Any], | |
| rows: list[list[Any]], | |
| output_root: Path, | |
| meta: dict[str, Any], | |
| ) -> None: | |
| model_name = model_display_name(model_spec) | |
| out_dir = output_root / model_output_slug(model_spec) | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| csv_path = out_dir / "all_results.csv" | |
| with csv_path.open("w", newline="") as handle: | |
| writer = csv.writer(handle) | |
| writer.writerow( | |
| [ | |
| "dataset", | |
| "model", | |
| "eval_metrics/MSE[mean]", | |
| "eval_metrics/MSE[0.5]", | |
| "eval_metrics/MAE[0.5]", | |
| "eval_metrics/MASE[0.5]", | |
| "eval_metrics/MAPE[0.5]", | |
| "eval_metrics/sMAPE[0.5]", | |
| "eval_metrics/MSIS", | |
| "eval_metrics/RMSE[mean]", | |
| "eval_metrics/NRMSE[mean]", | |
| "eval_metrics/ND[0.5]", | |
| "eval_metrics/mean_weighted_sum_quantile_loss", | |
| "domain", | |
| "num_variates", | |
| ] | |
| ) | |
| writer.writerows(rows) | |
| config = { | |
| "model": model_name, | |
| "model_type": model_spec.get("model_type", "zero-shot"), | |
| "model_dtype": "float32", | |
| "model_link": model_spec.get( | |
| "model_link", f"https://tsfm.ai/models/{model_spec['model_id']}" | |
| ), | |
| "code_link": "https://github.com/zhouziyu02/TSFM_Bench/blob/main/space/scripts/run_space_eval.py", | |
| "org": model_spec.get("org", "TSFM.ai"), | |
| "testdata_leakage": "No", | |
| "replication_code_available": "Yes", | |
| "api_model_id": model_spec["model_id"], | |
| } | |
| (out_dir / "config.json").write_text(json.dumps(config, indent=4) + "\n") | |
| (out_dir / "online_meta.json").write_text(json.dumps(meta, indent=4) + "\n") | |
| def write_dataset_properties(properties: dict[str, dict[str, Any]], output_root: Path) -> None: | |
| rows = [ | |
| { | |
| "dataset": key, | |
| "domain": values["domain"], | |
| "frequency": values["frequency"], | |
| "num_variates": values["num_variates"], | |
| } | |
| for key, values in sorted(properties.items()) | |
| ] | |
| csv_path = output_root / "dataset_properties.csv" | |
| with csv_path.open("w", newline="") as handle: | |
| writer = csv.DictWriter( | |
| handle, | |
| fieldnames=["dataset", "domain", "frequency", "num_variates"], | |
| ) | |
| writer.writeheader() | |
| writer.writerows(rows) | |
| def write_eval_state(payload: dict[str, Any]) -> None: | |
| OUTPUT_ROOT.mkdir(parents=True, exist_ok=True) | |
| EVAL_STATE_PATH.write_text(json.dumps(payload, indent=4) + "\n") | |
| def write_run_metadata(output_root: Path, payload: dict[str, Any]) -> None: | |
| (output_root / "online_status.json").write_text(json.dumps(payload, indent=4) + "\n") | |
| def make_predictor(model_spec: dict[str, Any], pred_len: int) -> TsfmApiPredictor: | |
| predictor = TsfmApiPredictor( | |
| config=TsfmApiConfig(model_id=model_spec["model_id"]), | |
| prediction_length=pred_len, | |
| quantile_levels=DEFAULT_QUANTILES, | |
| ) | |
| predictor.leaderboard_name = model_display_name(model_spec) | |
| return predictor | |
| def run_evaluation() -> int: | |
| if not os.getenv("TSFM_API_KEY"): | |
| write_eval_state( | |
| { | |
| "status": "disabled", | |
| "message": "TSFM_API_KEY not configured — set it in HF Space secrets.", | |
| "updated_at": datetime.now(timezone.utc).isoformat(), | |
| } | |
| ) | |
| logger.error("TSFM_API_KEY is missing; skipping evaluation") | |
| return 1 | |
| started = datetime.now(timezone.utc).isoformat() | |
| write_eval_state( | |
| { | |
| "status": "running", | |
| "started_at": started, | |
| "message": "Collecting TS-Bench data and calling TSFM.ai API…", | |
| "updated_at": started, | |
| } | |
| ) | |
| try: | |
| source = load_data_source(DATA_CONFIG) | |
| if isinstance(source, TsBenchDataSource) and source._settings.auto_refresh: | |
| source.refresh() | |
| datasets = source.list_datasets() | |
| if not datasets: | |
| raise RuntimeError("No TS-Bench tasks available after data collection") | |
| properties = load_dataset_properties(DATA_CONFIG) | |
| model_specs = load_model_specs(MODEL_CONFIG) | |
| OUTPUT_ROOT.mkdir(parents=True, exist_ok=True) | |
| write_dataset_properties(properties, OUTPUT_ROOT) | |
| failed_models: list[str] = [] | |
| all_model_meta: dict[str, Any] = {} | |
| for model_spec in model_specs: | |
| model_name = model_display_name(model_spec) | |
| rows: list[list[Any]] = [] | |
| dataset_meta: list[dict[str, Any]] = [] | |
| try: | |
| for ds_name in datasets: | |
| pred_len = source.get_prediction_length(ds_name) | |
| predictor = make_predictor(model_spec, pred_len) | |
| try: | |
| logger.info("Evaluating %s on %s", model_name, ds_name) | |
| result = run_online_eval_for_dataset(source, ds_name, predictor) | |
| except Exception: | |
| logger.exception("Skipping %s on %s", model_name, ds_name) | |
| continue | |
| rows.append( | |
| [ | |
| result.dataset, | |
| model_name, | |
| result.metrics["MSE[mean]"], | |
| result.metrics["MSE[0.5]"], | |
| result.metrics["MAE[0.5]"], | |
| result.metrics["MASE[0.5]"], | |
| result.metrics["MAPE[0.5]"], | |
| result.metrics["sMAPE[0.5]"], | |
| result.metrics["MSIS"], | |
| result.metrics["RMSE[mean]"], | |
| result.metrics["NRMSE[mean]"], | |
| result.metrics["ND[0.5]"], | |
| result.metrics["mean_weighted_sum_quantile_loss"], | |
| result.domain, | |
| result.num_variates, | |
| ] | |
| ) | |
| dataset_meta.append( | |
| { | |
| "dataset": ds_name, | |
| "data_fetched_at": result.data_fetched_at, | |
| "context_length": result.context_length, | |
| "prediction_length": result.prediction_length, | |
| } | |
| ) | |
| if not rows: | |
| raise RuntimeError(f"No successful datasets for {model_name}") | |
| model_meta = { | |
| "model": model_name, | |
| "api_model_id": model_spec["model_id"], | |
| "evaluated_at": datetime.now(timezone.utc).isoformat(), | |
| "datasets": dataset_meta, | |
| } | |
| write_model_results(model_spec, rows, OUTPUT_ROOT, model_meta) | |
| all_model_meta[model_name] = model_meta | |
| except Exception: | |
| logger.exception("Failed evaluating %s", model_name) | |
| failed_models.append(model_name) | |
| finished = datetime.now(timezone.utc).isoformat() | |
| status = "ok" if not failed_models else "partial" | |
| write_run_metadata( | |
| OUTPUT_ROOT, | |
| { | |
| "status": status, | |
| "started_at": started, | |
| "finished_at": finished, | |
| "data_source": "ts_bench", | |
| "data_config": str(DATA_CONFIG), | |
| "ts_bench_root": os.environ.get("TS_BENCH_ROOT", ""), | |
| "models": all_model_meta, | |
| "failed_models": failed_models, | |
| }, | |
| ) | |
| write_eval_state( | |
| with_next_eval_fields( | |
| { | |
| "status": status, | |
| "started_at": started, | |
| "finished_at": finished, | |
| "failed_models": failed_models, | |
| "task_count": len(datasets), | |
| "model_count": len(model_specs) - len(failed_models), | |
| "message": "Evaluation complete", | |
| "updated_at": finished, | |
| }, | |
| finished, | |
| BENCHMARK_INTERVAL_SECONDS, | |
| ) | |
| ) | |
| return 0 if status == "ok" else 2 | |
| except Exception as exc: | |
| finished = datetime.now(timezone.utc).isoformat() | |
| logger.exception("Evaluation run failed") | |
| write_eval_state( | |
| with_next_eval_fields( | |
| { | |
| "status": "error", | |
| "started_at": started, | |
| "finished_at": finished, | |
| "message": str(exc), | |
| "updated_at": finished, | |
| }, | |
| finished, | |
| BENCHMARK_INTERVAL_SECONDS, | |
| ) | |
| ) | |
| return 1 | |
| if __name__ == "__main__": | |
| raise SystemExit(run_evaluation()) | |