File size: 4,827 Bytes
b560af6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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)