#!/usr/bin/env python3 """Periodically refresh TS-Bench data, run TSFM.ai zero-shot eval, push to HF Space. Supports per-dataset eval intervals via source_intervals in the data config. Each tick checks which datasets are due for evaluation and only runs those. """ from __future__ import annotations import argparse import fcntl import json import logging import math import os import signal import subprocess import sys import time from datetime import datetime, timezone from pathlib import Path import yaml from dotenv import load_dotenv REPO_ROOT = Path(__file__).resolve().parents[1] load_dotenv(REPO_ROOT / ".env") DEFAULT_INTERVAL_SECONDS = int(os.getenv("TSFM_BENCH_INTERVAL_SECONDS", "300")) PUSH_RETRIES = int(os.getenv("TSFM_PUSH_RETRIES", "3")) PUSH_MIN_INTERVAL_SECONDS = max( 0, int(os.getenv("TSFM_HF_PUSH_MIN_INTERVAL_SECONDS", "300")) ) MAX_DATASETS_PER_CYCLE = max( 1, int(os.getenv("TSFM_MAX_DATASETS_PER_CYCLE", "4")) ) LOCK_FILE = REPO_ROOT / ".online_daemon.lock" _last_successful_push_at = 0.0 REMOTE_STATE = None # Set only by cloud/worker.py after a verified restore. logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") logger = logging.getLogger(__name__) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( "--interval-minutes", type=int, default=None, help="Minutes between cycles (overrides --interval-seconds). " "Used as the global fallback when source_intervals is not set for a dataset.", ) parser.add_argument( "--interval-seconds", type=int, default=DEFAULT_INTERVAL_SECONDS, help="Seconds between cycles (default: TSFM_BENCH_INTERVAL_SECONDS or 300). " "Used as the global fallback when source_intervals is not set for a dataset.", ) parser.add_argument( "--data-config", type=Path, default=Path("configs/datasets/ts_bench.yaml"), ) parser.add_argument( "--model-config", type=Path, default=Path("configs/models/online_tsfm.yaml"), ) parser.add_argument( "--output-root", type=Path, default=Path("space/results"), ) parser.add_argument( "--once", action="store_true", help="Run a single collect + eval + push cycle and exit (evaluates all datasets)", ) parser.add_argument( "--no-push", action="store_true", help="Skip pushing results to Hugging Face", ) return parser.parse_args() def global_interval(args: argparse.Namespace) -> int: if args.interval_minutes is not None: return max(30, args.interval_minutes * 60) return max(30, args.interval_seconds) class CycleLock: """Prevent overlapping cycles when eval takes longer than the interval.""" def __init__(self, path: Path): self._path = path self._handle = None def __enter__(self) -> bool: self._path.parent.mkdir(parents=True, exist_ok=True) self._handle = self._path.open("w") try: fcntl.flock(self._handle.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) except BlockingIOError: logger.warning("Previous cycle still running; skipping this tick") self._handle.close() self._handle = None return False self._handle.write(f"pid={os.getpid()} started={datetime.now(timezone.utc).isoformat()}\n") self._handle.flush() return True def __exit__(self, exc_type, exc, tb) -> None: if self._handle is not None: fcntl.flock(self._handle.fileno(), fcntl.LOCK_UN) self._handle.close() # --------------------------------------------------------------------------- # Per-dataset interval scheduling # --------------------------------------------------------------------------- def load_source_intervals(data_config: Path) -> dict[str, int]: """Read source_intervals from the data config YAML.""" if not data_config.exists(): return {} try: payload = yaml.safe_load(data_config.read_text()) except Exception: return {} raw = payload.get("source_intervals", {}) return {str(k): int(v) for k, v in raw.items()} def load_dataset_intervals(data_config: Path, fallback: int, output_root: Path | None = None) -> dict[str, int]: """Load task IDs and their per-dataset eval intervals from the data source. Returns {task_id: interval_seconds}. Uses TsBenchDataSource.get_eval_interval() which resolves via source_id matching against source_intervals config. """ sys.path.insert(0, str(REPO_ROOT / "src")) from tsfm_bench.data.registry import load_data_source from tsfm_bench.data.ts_bench import TsBenchDataSource source = load_data_source(data_config) task_ids = list(source.list_datasets()) if isinstance(source, TsBenchDataSource): intervals = {tid: source.get_eval_interval(tid, fallback) for tid in task_ids} if output_root: try: mapping = {source._tasks[tid].leaderboard_name: source.get_eval_interval(tid, fallback) for tid in task_ids} output_root.mkdir(parents=True, exist_ok=True) (output_root / "dataset_eval_intervals.json").write_text(json.dumps(mapping, indent=4) + "\n") except Exception as e: logger.warning("Could not save dataset_eval_intervals.json: %s", e) return intervals # Non-TsBench sources: use uniform fallback intervals = {tid: fallback for tid in task_ids} if output_root: try: mapping = {tid: fallback for tid in task_ids} output_root.mkdir(parents=True, exist_ok=True) (output_root / "dataset_eval_intervals.json").write_text(json.dumps(mapping, indent=4) + "\n") except Exception as e: logger.warning("Could not save dataset_eval_intervals.json: %s", e) return intervals def compute_tick_interval(intervals: dict[str, int], fallback: int) -> int: """Compute the GCD of all configured intervals as the daemon tick.""" all_values = list(intervals.values()) or [fallback] result = all_values[0] for v in all_values[1:]: result = math.gcd(result, v) return max(10, result) # Floor at 10s to avoid spinning def clean_task_id(tid: str) -> str: import re # Strip suffix like _20260613t130000z return re.sub(r'_\d{8}t\d{6}z$', '', tid, flags=re.IGNORECASE) def get_due_datasets( last_eval: dict[str, float], dataset_intervals: dict[str, int], now: float, ) -> list[str]: """Return task IDs that are due for evaluation based on their intervals.""" due = [] for tid, interval in dataset_intervals.items(): base_tid = clean_task_id(tid) elapsed = now - last_eval.get(base_tid, 0.0) if elapsed >= interval: due.append(tid) return due def select_due_batch(due: list[str], last_eval: dict[str, float]) -> list[str]: """Bound cycle size while giving the stalest datasets priority.""" ordered = sorted( due, key=lambda tid: (last_eval.get(clean_task_id(tid), 0.0), clean_task_id(tid)), ) return ordered[:MAX_DATASETS_PER_CYCLE] # --------------------------------------------------------------------------- # Eval + push # --------------------------------------------------------------------------- def run_eval(args: argparse.Namespace, datasets: list[str] | None = None) -> int: cmd = [ sys.executable, str(REPO_ROOT / "scripts" / "run_online_eval.py"), "--data-config", str(args.data_config), "--model-config", str(args.model_config), "--output-root", str(args.output_root), ] if datasets: cmd += ["--datasets"] + datasets logger.info("Running: %s", " ".join(cmd)) return run_managed_subprocess( cmd, timeout=1800, timeout_message="Evaluation task timed out after 30 minutes", ) def subprocess_env() -> dict[str, str]: env = os.environ.copy() src_path = str(REPO_ROOT / "src") pythonpath = env.get("PYTHONPATH") env["PYTHONPATH"] = src_path if not pythonpath else os.pathsep.join([src_path, pythonpath]) return env def run_managed_subprocess( cmd: list[str], *, timeout: int, timeout_message: str, ) -> int: """Run a child in its own process group and reap all children on timeout.""" process = subprocess.Popen( cmd, cwd=str(REPO_ROOT), env=subprocess_env(), start_new_session=True, ) try: return process.wait(timeout=timeout) except (KeyboardInterrupt, SystemExit): try: os.killpg(process.pid, signal.SIGTERM) process.wait(timeout=10) except subprocess.TimeoutExpired: os.killpg(process.pid, signal.SIGKILL) process.wait() except ProcessLookupError: pass raise except subprocess.TimeoutExpired: logger.error(timeout_message) try: os.killpg(process.pid, signal.SIGTERM) except ProcessLookupError: pass try: process.wait(timeout=10) except subprocess.TimeoutExpired: try: os.killpg(process.pid, signal.SIGKILL) except ProcessLookupError: pass process.wait() return 124 def load_eval_status(output_root: Path) -> dict[str, object]: status_path = output_root / "online_status.json" try: payload = json.loads(status_path.read_text()) except (OSError, json.JSONDecodeError): return {} return payload if isinstance(payload, dict) else {} def eval_status_is_publishable(payload: dict[str, object]) -> bool: try: task_count = int(payload.get("task_count") or 0) except (TypeError, ValueError): return False return ( payload.get("status") in {"ok", "partial"} and payload.get("aggregate_status", "ok") == "ok" and task_count > 0 ) def push_is_due(now: float | None = None) -> bool: if PUSH_MIN_INTERVAL_SECONDS == 0 or _last_successful_push_at <= 0: return True current = time.time() if now is None else now return current - _last_successful_push_at >= PUSH_MIN_INTERVAL_SECONDS def push_results(output_root: Path) -> int: cmd = [sys.executable, str(REPO_ROOT / "scripts" / "push_results_to_hf.py"), "--results", str(output_root)] logger.info("Pushing results to HF Space") return run_managed_subprocess( cmd, timeout=600, timeout_message="Hugging Face push task timed out after 10 minutes", ) def push_with_retries(output_root: Path) -> int: last_code = 1 for attempt in range(1, PUSH_RETRIES + 1): code = push_results(output_root) if code == 0: return 0 last_code = code if attempt < PUSH_RETRIES: wait = 10 * attempt logger.warning( "HF push attempt %s/%s failed; retrying in %ss", attempt, PUSH_RETRIES, wait, ) time.sleep(wait) return last_code def run_cycle(args: argparse.Namespace, datasets: list[str] | None = None) -> int: global _last_successful_push_at with CycleLock(LOCK_FILE) as acquired: if not acquired: return 0 started = datetime.now(timezone.utc).isoformat() if datasets: logger.info("Cycle started at %s — evaluating %d datasets: %s", started, len(datasets), datasets) else: logger.info("Cycle started at %s — evaluating all datasets", started) eval_code = run_eval(args, datasets) if REMOTE_STATE is not None: # Preserve newly frozen predictions even after partial/failed cycles. # Publishing is forbidden until the checkpoint commit succeeds. REMOTE_STATE.checkpoint() if eval_code != 0: logger.error("Evaluation failed with exit code %s", eval_code) eval_status = load_eval_status(args.output_root) if eval_code == 0 and not eval_status_is_publishable(eval_status): logger.error( "Evaluation produced no publishable tasks (status=%s, task_count=%s)", eval_status.get("status", "missing"), eval_status.get("task_count", "missing"), ) eval_code = 3 if args.no_push: logger.info("Cycle complete (push skipped)") return eval_code if eval_code != 0: logger.warning("Skipping HF push because evaluation was not publishable") return eval_code if not args.once and not push_is_due(): remaining = max( 1, int(PUSH_MIN_INTERVAL_SECONDS - (time.time() - _last_successful_push_at)), ) logger.info("HF push deferred for %ss to coalesce rapid updates", remaining) return 0 push_code = push_with_retries(args.output_root) if push_code != 0: logger.error("HF push failed after %s attempts", PUSH_RETRIES) return push_code _last_successful_push_at = time.time() logger.info("Cycle complete: eval + push OK") return 0 def load_last_eval_from_history(history_path: Path, fallback: int, data_config: Path) -> dict[str, float]: """Read eval_history.jsonl and initialize last_eval timestamps for task IDs.""" last_eval = {} if not history_path.exists(): return last_eval latest_by_leaderboard = {} try: with open(history_path, "r", encoding="utf-8") as f: for line in f: if not line.strip(): continue try: entry = json.loads(line) ds = entry.get("dataset") eval_at = entry.get("evaluated_at") if ds and eval_at: # Parse ISO format (handling Z or offset) dt = datetime.fromisoformat(eval_at.replace("Z", "+00:00")) ts = dt.timestamp() if ds not in latest_by_leaderboard or ts > latest_by_leaderboard[ds]: latest_by_leaderboard[ds] = ts except Exception: pass except Exception as e: logger.warning("Failed to read eval_history.jsonl for initial scheduling: %s", e) # Resolve source to map leaderboard names back to clean task IDs try: sys.path.insert(0, str(REPO_ROOT / "src")) from tsfm_bench.data.registry import load_data_source from tsfm_bench.data.ts_bench import TsBenchDataSource source = load_data_source(data_config) if isinstance(source, TsBenchDataSource): for tid in source.list_datasets(): try: task = source._tasks[tid] base_tid = clean_task_id(tid) if task.leaderboard_name in latest_by_leaderboard: last_eval[base_tid] = latest_by_leaderboard[task.leaderboard_name] logger.info("Initialized last_eval for %s: %s", base_tid, datetime.fromtimestamp(last_eval[base_tid], tz=timezone.utc).isoformat()) except Exception: pass except Exception as e: logger.warning("Could not load datasource for initializing last_eval: %s", e) return last_eval def main() -> None: signal.signal(signal.SIGTERM, lambda _signum, _frame: sys.exit(0)) global _last_successful_push_at args = parse_args() fallback = global_interval(args) previous_status = load_eval_status(args.output_root) pushed_at = previous_status.get("pushed_at") if isinstance(pushed_at, str): try: _last_successful_push_at = datetime.fromisoformat( pushed_at.replace("Z", "+00:00") ).timestamp() except ValueError: pass # --once: evaluate all datasets in a single cycle, then exit if args.once: raise SystemExit(run_cycle(args)) # Load per-dataset intervals via data source (uses source_id matching) dataset_intervals: dict[str, int] = {} try: dataset_intervals = load_dataset_intervals(args.data_config, fallback, args.output_root) logger.info("Discovered %d dataset tasks with per-dataset intervals:", len(dataset_intervals)) for tid, interval in sorted(dataset_intervals.items()): logger.info(" %s → every %ds", tid, interval) except Exception: logger.warning("Could not load dataset intervals; falling back to uniform scheduling") has_varied_intervals = len(set(dataset_intervals.values())) > 1 if has_varied_intervals: tick = compute_tick_interval(dataset_intervals, fallback) logger.info("Per-dataset scheduling: tick=%ds (GCD of all intervals)", tick) else: tick = fallback logger.info("Uniform scheduling: interval=%ds", tick) # Initialize last_eval from evaluation history files history_path = args.output_root / "eval_history.jsonl" last_eval = load_last_eval_from_history(history_path, fallback, args.data_config) while True: try: # Reload intervals to discover newly pulled task IDs try: dataset_intervals = load_dataset_intervals(args.data_config, fallback, args.output_root) except Exception: pass now = time.time() if has_varied_intervals and dataset_intervals: due = get_due_datasets(last_eval, dataset_intervals, now) if due: due_count = len(due) due = select_due_batch(due, last_eval) logger.info( "Due datasets (%d/%d; running stalest %d): %s", due_count, len(dataset_intervals), len(due), due, ) cycle_code = run_cycle(args, due) eval_time = time.time() # Load failed datasets from online_status.json to implement smart retry failed_tids = set() status_path = args.output_root / "online_status.json" if status_path.exists(): try: status_data = json.loads(status_path.read_text()) failed_tids = set(status_data.get("failed_datasets", [])) except Exception: pass if cycle_code != 0 and not failed_tids: failed_tids.update(due) for tid in due: base_tid = clean_task_id(tid) if tid in failed_tids or base_tid in failed_tids: # Failed: retry in 5 minutes interval = dataset_intervals.get(tid, fallback) last_eval[base_tid] = eval_time - interval + 300 logger.warning("Dataset %s failed evaluation; scheduled for retry in 5 minutes", tid) else: last_eval[base_tid] = eval_time else: logger.debug("No datasets due this tick") else: # Uniform mode: evaluate all datasets run_cycle(args) eval_time = time.time() for tid in dataset_intervals: last_eval[clean_task_id(tid)] = eval_time except Exception as err: logger.exception("Unexpected error in daemon loop iteration: %s", err) if REMOTE_STATE is not None: # The supervisor restores the winning remote revision before a # retry; a stale writer must never overwrite another checkpoint. raise if tick >= 60 and tick % 60 == 0: logger.info("Sleeping %s minutes until next tick", tick // 60) else: logger.info("Sleeping %s seconds until next tick", tick) time.sleep(tick) if __name__ == "__main__": main()