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