gollem-v5-ckpts / fabryka_push.py
Maggio33's picture
Upload fabryka_push.py with huggingface_hub
b560af6 verified
Raw History Blame
4.83 kB
#!/usr/bin/env python3
"""Push a trainer's events.jsonl to track.fabryka.ai via the Fabryka SDK (api-key auth).
Reads the events.jsonl the trainer emits (--events-jsonl), backfills all existing 'update'
records, then tails new ones every 15s. Logs metrics (loss, tokens_per_second, gradient_norm,
learning_rate) at step=updates. Read-only on the trainer; no GPU, no weights.
Usage: FABRYKA_API_KEY=... python fabryka_push.py EVENTS.jsonl RUN_NAME [PROJECT]
"""
import json
import os
import sys
import time
def _ensure_sdk_env_support():
"""Old installed fabryka SDKs ignore FABRYKA_RUN_ID / FABRYKA_PUBLIC_LIVE_TRACKING.
Patch the installed client.py so a stable run id + public live tracking work."""
try:
import fabryka
p = os.path.join(os.path.dirname(fabryka.__file__), "client.py")
s = open(p, encoding="utf-8").read()
orig = s
if "FABRYKA_RUN_ID" not in s:
s = s.replace("self.run_id = str(uuid.uuid4())",
'self.run_id = os.getenv("FABRYKA_RUN_ID") or str(uuid.uuid4())')
if "public_live_tracking" not in s:
anchor = 'commit = _command("git", "rev-parse", "HEAD")'
if anchor in s:
inject = ('if os.getenv("FABRYKA_PUBLIC_LIVE_TRACKING") == "1":\n'
' metadata["public_live_tracking"] = True\n ')
s = s.replace(anchor, inject + anchor, 1)
if s != orig:
open(p, "w", encoding="utf-8").write(s)
# drop cached (unpatched) module so the later `from fabryka import` reloads patched code
for mod in [m for m in sys.modules if m == "fabryka" or m.startswith("fabryka.")]:
del sys.modules[mod]
print("[fabryka_push] patched installed SDK for env support", flush=True)
except Exception as e:
print(f"[fabryka_push] SDK patch skipped: {e}", flush=True)
_ensure_sdk_env_support()
from fabryka import RunClient
events_path = sys.argv[1]
name = sys.argv[2] if len(sys.argv) > 2 else "run"
project = sys.argv[3] if len(sys.argv) > 3 else "gollem-v5"
total_tokens = int(sys.argv[4]) if len(sys.argv) > 4 else 0 # planned tokens -> enables ETA + progress
client = RunClient(api_url=os.environ.get("FABRYKA_API_URL", "https://track.fabryka.ai"),
api_key=os.environ["FABRYKA_API_KEY"])
client.init(project=project, name=name)
print(f"[fabryka_push] init project={project} name={name} events={events_path}", flush=True)
seen = 0
idle = 0
while True:
try:
with open(events_path, encoding="utf-8") as f:
lines = f.readlines()
except FileNotFoundError:
time.sleep(5)
continue
new = lines[seen:]
seen = len(lines)
pushed = 0
end = False
for ln in new:
try:
e = json.loads(ln)
except Exception:
continue
kind = e.get("kind")
if kind == "update" and e.get("metrics"):
raw = e["metrics"]
toks = e.get("tokens", 0)
tps = raw.get("tokens_per_second", 0)
m = {}
# whitelisted keys -> visible to non-owner viewers (Kacper); see api.py run_detail
if "loss" in raw: m["train/loss"] = raw["loss"]
if tps: m["throughput/tokens_sec"] = tps
if toks: m["training/tokens_seen"] = toks
if total_tokens and toks: m["progress"] = round(100.0 * toks / total_tokens, 2)
# extras -> owner (Arek) sees these too
for k in ("gnorm", "lr", "muon_lr"):
if k in raw: m[k] = raw[k]
if toks: m["tokens_b"] = round(toks / 1e9, 3)
if total_tokens and toks and tps:
m["eta_hours"] = round(max(0, total_tokens - toks) / tps / 3600.0, 2)
m = {k: v for k, v in m.items() if isinstance(v, (int, float))}
if m:
client.log(m, step=int(e.get("updates", 0)))
pushed += 1
elif kind == "evaluation" and e.get("metrics"):
ev = e["metrics"]
em = {f"eval_{k}": v for k, v in ev.items() if isinstance(v, (int, float))}
# whitelisted val keys -> non-owner sees eval trajectory
if "byte_ppl" in ev: em["val/perplexity"] = ev["byte_ppl"]
if "val_loss" in ev: em["val/loss"] = ev["val_loss"]
if em:
client.log(em, step=int(e.get("updates", 0)))
pushed += 1
elif kind in ("end", "failed"):
end = True
if pushed:
print(f"[fabryka_push] pushed {pushed} updates (total lines={seen})", flush=True)
idle = 0
else:
idle += 1
if end:
print("[fabryka_push] end event seen; finishing", flush=True)
client.finish()
break
time.sleep(15)