Spaces:
Running on Zero
Running on Zero
File size: 6,458 Bytes
8028640 | 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 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """Enroll the v1 model line-up and write `arena/registry.json`.
Every model goes through the same `runtime.enroll` path a user would: the
revision is pinned to an immutable sha, a smoke test runs, and the result is
recorded. Nothing is registered that has not actually produced a forecast.
Latency is measured here too, and a CPU-tier model that misses its budget is
demoted to GPU tier in the registry rather than left to time out in front of
someone. The measurement is only meaningful on the hardware it was taken on,
so the machine is recorded alongside the number -- run this on the Space to
get numbers that describe the Space.
HF_TOKEN=$(cat ../../key.txt) python scripts/bootstrap_registry.py
python scripts/bootstrap_registry.py --dry-run # no writes
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import platform
import sys
import time
from pathlib import Path
import numpy as np
import pandas as pd
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from src import config, runtime # noqa: E402
from src.adapters import get_adapter # noqa: E402
from src.store import ArenaStore, now_utc # noqa: E402
logging.basicConfig(level=logging.INFO, format="%(levelname)-7s %(message)s")
log = logging.getLogger("bootstrap")
# The v1 line-up. Order is the order they appear in the UI.
SEED_MODELS = [
("baseline", "baseline/random-walk"),
("baseline", "baseline/drift"),
("baseline", "baseline/bootstrap"),
("chronos", "amazon/chronos-bolt-tiny"),
("chronos", "amazon/chronos-bolt-mini"),
("chronos", "amazon/chronos-bolt-small"),
("chronos", "amazon/chronos-bolt-base"),
("timesfm", "google/timesfm-2.5-200m-pytorch"),
("kronos", "NeoQuasar/Kronos-mini"),
("kronos", "NeoQuasar/Kronos-small"),
("kronos", "NeoQuasar/Kronos-base"),
]
# The forecast the budget is defined against: a default 1h forecast.
BUDGET_HORIZON = 24
BUDGET_SAMPLES = config.DEFAULT_N_SAMPLES
def _synthetic(n=config.DEFAULT_CONTEXT_BARS):
rng = np.random.default_rng(7)
close = 50000.0 * np.exp(np.cumsum(rng.normal(0, 0.01, n)))
openp = np.concatenate([[50000.0], close[:-1]])
return pd.DataFrame({
"ts": pd.date_range("2025-01-01", periods=n, freq="1h", tz="UTC"),
"open": openp,
"high": np.maximum(openp, close) * 1.002,
"low": np.minimum(openp, close) * 0.998,
"close": close,
"volume": np.abs(rng.normal(1000, 200, n)) + 1.0,
})
def measure(family: str, model_id: str, context) -> dict:
"""Cold and warm latency for a default forecast."""
adapter = get_adapter(family, model_id)
# Measure on the context the app actually sends, not on the model's own
# ceiling: the cost is roughly linear in context length, so measuring at
# `max_context` would report a number no user ever waits for.
context = context.iloc[-min(len(context), config.DEFAULT_CONTEXT_BARS):]
t0 = time.time()
adapter.load()
load_s = time.time() - t0
t0 = time.time()
adapter.predict(context, horizon=BUDGET_HORIZON,
n_samples=BUDGET_SAMPLES, seed=0)
cold_s = time.time() - t0
warm = []
for _ in range(3):
t0 = time.time()
adapter.predict(context, horizon=BUDGET_HORIZON,
n_samples=BUDGET_SAMPLES, seed=0)
warm.append(time.time() - t0)
warm_s = float(np.median(warm))
declared = adapter.capabilities().hardware
# "Cold" for a user is load plus first forecast: the weights are not in
# memory when they arrive.
total_cold = load_s + cold_s
misses = declared == "cpu" and (
total_cold > config.CPU_BUDGET_COLD_S or warm_s > config.CPU_BUDGET_WARM_S)
return {
"declared_hardware": declared,
"hardware": "gpu" if misses else declared,
"demoted": bool(misses),
"load_s": round(load_s, 3),
"cold_s": round(total_cold, 3),
"warm_s": round(warm_s, 3),
"horizon": BUDGET_HORIZON,
"n_samples": BUDGET_SAMPLES,
"machine": f"{platform.system()}-{platform.machine()}-py{platform.python_version()}",
"on_space": bool(os.environ.get("SPACE_ID")),
"measured_ts": now_utc().isoformat(),
}
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--dry-run", action="store_true",
help="measure and print, write nothing")
ap.add_argument("--skip-measure", action="store_true",
help="enroll without a latency run (declared tiers stand)")
ap.add_argument("--local-root", default=None)
args = ap.parse_args()
offline = args.dry_run or not os.environ.get("HF_TOKEN")
store = ArenaStore(local_root=args.local_root, offline=offline)
registry = store.get_registry()
context = _synthetic()
for family, model_id in SEED_MODELS:
log.info("enrolling %s (%s)", model_id, family)
outcome = runtime.enroll(store, family, model_id,
enrolled_by="bit-trading-company",
registry=registry)
if not outcome.ok:
log.error(" refused: %s", outcome.message)
continue
if outcome.already:
log.info(" already enrolled at this revision")
entry = registry["models"][outcome.model_slug]
if not args.skip_measure:
try:
latency = measure(family, model_id, context)
entry["latency"] = latency
entry["capabilities"]["hardware"] = latency["hardware"]
flag = " DEMOTED to gpu" if latency["demoted"] else ""
log.info(" cold=%.1fs warm=%.2fs tier=%s%s",
latency["cold_s"], latency["warm_s"],
latency["hardware"], flag)
except Exception as e:
log.error(" latency run failed: %s", e)
entry["latency"] = {"error": str(e)[:200]}
registry["updated_ts"] = now_utc().isoformat()
registry["version"] = 1
if args.dry_run:
print(json.dumps(registry, indent=2)[:4000])
log.info("dry run: nothing written")
return 0
store.put_registry(registry)
pushed = store.flush("arena: bootstrap model registry")
log.info("registry written (%d file(s) pushed)", pushed)
return 0
if __name__ == "__main__":
raise SystemExit(main())
|