LiveHouse-TS / scripts /run_online_daemon.py
ziyuzhou02's picture
Deploy GitHub main 3feb6cda1511
e317359 verified
Raw History Blame Contribute Delete
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()