File size: 10,405 Bytes
4fd620e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
177
178
179
180
181
182
183
184
185
186
"""SpikeWhale training control panel (CLI: python -m daisychain.spikewhale_panel).

A web page with sliders for the SpikeWhale config. Pick a size your hardware can
handle, hit Start, and it launches the real DaisyChain training (SpikeWhale +
FineWeb-Edu) and streams the live loss. The exact env is shown so you can run the
same command on other machines to train distributed.
"""
import json
import os
import subprocess
import sys
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer

PORT = int(os.environ.get("SW_PANEL_PORT", "8899"))
_proc = None
_log = []          # rolling training log
_lock = threading.Lock()


def _pump(proc):
    for line in iter(proc.stdout.readline, ""):
        with _lock:
            _log.append(line.rstrip("\n"))
            if len(_log) > 400:
                del _log[:200]
    proc.stdout.close()


def start_training(cfg):
    global _proc
    if _proc and _proc.poll() is None:
        return False, "already running"
    env = dict(os.environ)
    env.update({
        "MASTER_ADDR": "127.0.0.1", "MASTER_PORT": "29610", "WORLD_SIZE": "1",
        "RANK": "0", "USE_LIBUV": "0", "PYTHONUNBUFFERED": "1",
        "DAISY_TASK": "daisychain.spikewhale_task:SpikeWhaleTask",
        "DAISY_SW_HIDDEN": str(cfg["hidden"]), "DAISY_SW_LAYERS": str(cfg["layers"]),
        "DAISY_SW_HEADS": str(cfg["heads"]), "DAISY_SW_EXPERTS": str(cfg["experts"]),
        "DAISY_SW_SEQLEN": str(cfg["seqlen"]),
        "DAISY_STEPS": str(cfg["steps"]), "DAISY_LR": str(cfg["lr"]),
        "DAISY_OPTIMIZER": "adam", "DAISY_BASE_BATCH": str(cfg["batch"]),
    })
    if cfg.get("dataset"):
        env["DAISY_SW_DATASET"] = cfg["dataset"]
        # blank subset means the dataset's default config; unset any inherited one
        if cfg.get("subset"):
            env["DAISY_SW_SUBSET"] = cfg["subset"]
        else:
            env.pop("DAISY_SW_SUBSET", None)
    with _lock:
        _log.clear()
        _log.append(f"launching training: hidden={cfg['hidden']} layers={cfg['layers']} "
                    f"experts={cfg['experts']} seqlen={cfg['seqlen']} steps={cfg['steps']} "
                    f"dataset={env.get('DAISY_SW_DATASET', 'HuggingFaceFW/fineweb-edu')}"
                    + (f":{env['DAISY_SW_SUBSET']}" if env.get("DAISY_SW_SUBSET") else ""))
    _proc = subprocess.Popen([sys.executable, "-u", "-m", "daisychain.train"],
                             env=env, stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
                             text=True, bufsize=1)
    threading.Thread(target=_pump, args=(_proc,), daemon=True).start()
    return True, "started"


PAGE = """<!doctype html><html><head><meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>SpikeWhale · DaisyChain trainer</title>
<style>
 body{font-family:system-ui,-apple-system,Segoe UI,Roboto,sans-serif;max-width:720px;margin:0 auto;
      padding:22px;background:#efe4c9;color:#2a1d0a}
 @media(prefers-color-scheme:dark){body{background:#14100a;color:#ede1c3}}
 h1{margin:0 0 2px}.sub{color:#6b4423;margin:0 0 16px}
 @media(prefers-color-scheme:dark){.sub{color:#c9b072}}
 .card{background:#fbf6e8;border:1px solid rgba(139,111,71,.3);border-radius:10px;padding:16px;margin:12px 0}
 @media(prefers-color-scheme:dark){.card{background:#1f1a12;border-color:rgba(201,176,114,.35)}}
 label{display:flex;justify-content:space-between;font-size:.92rem;margin:10px 0 4px;font-weight:600}
 input[type=range]{width:100%;accent-color:#4a7c2e}
 input[type=text]{width:100%;padding:8px 10px;border-radius:8px;border:1px solid rgba(139,111,71,.4);
        background:rgba(255,255,255,.5);color:inherit;font-family:'Courier New',monospace;font-size:.9rem}
 @media(prefers-color-scheme:dark){input[type=text]{background:rgba(0,0,0,.25);border-color:rgba(201,176,114,.35)}}
 .val{font-family:'Courier New',monospace;color:#4a7c2e;font-weight:700}
 @media(prefers-color-scheme:dark){.val{color:#9bc466}}
 button{background:linear-gradient(135deg,#4a7c2e,#2d5016);color:#f5ecd9;border:0;border-radius:8px;
        padding:12px 26px;font-weight:800;font-size:1rem;cursor:pointer}
 .num{font-family:'Courier New',monospace;font-size:2rem;font-weight:700;text-align:center;
      color:#f5ecd9;background:linear-gradient(135deg,#2d5016,#1f3a0f);border-radius:10px;padding:14px}
 pre{background:rgba(0,0,0,.06);border-radius:8px;padding:10px;max-height:220px;overflow:auto;
     font-size:.78rem;white-space:pre-wrap;font-family:'Courier New',monospace}
 @media(prefers-color-scheme:dark){pre{background:rgba(0,0,0,.3)}}
 .lbl{font-size:11px;font-weight:800;letter-spacing:1.5px;text-transform:uppercase;color:#6b4423;margin-bottom:8px}
 @media(prefers-color-scheme:dark){.lbl{color:#c9b072}}
</style></head><body>
<h1>🐋 SpikeWhale · DaisyChain</h1>
<p class="sub">Pick a size your hardware can train, then start. Trains the real SpikeWhale on streamed FineWeb-Edu, distributed by DaisyChain. Smaller = faster on old hardware.</p>
<div id="settings">
<div class="card"><div class="lbl">Model size</div>
 <label>Hidden size <span class="val" id="vhidden">256</span></label><input type="range" id="hidden" min="64" max="768" step="64" value="256">
 <label>Layers <span class="val" id="vlayers">4</span></label><input type="range" id="layers" min="1" max="12" step="1" value="4">
 <label>Attention heads <span class="val" id="vheads">4</span></label><input type="range" id="heads" min="1" max="8" step="1" value="4">
 <label>MoE experts <span class="val" id="vexperts">4</span></label><input type="range" id="experts" min="1" max="8" step="1" value="4">
 <label>Sequence length <span class="val" id="vseqlen">128</span></label><input type="range" id="seqlen" min="32" max="512" step="32" value="128">
</div>
<div class="card"><div class="lbl">Training</div>
 <label>Learning rate ×1e-4 <span class="val" id="vlr">30</span></label><input type="range" id="lr" min="1" max="100" step="1" value="30">
 <label>Batch (per step) <span class="val" id="vbatch">4</span></label><input type="range" id="batch" min="1" max="16" step="1" value="4">
 <label>Steps <span class="val" id="vsteps">200</span></label><input type="range" id="steps" min="20" max="2000" step="20" value="200">
</div>
<div class="card"><div class="lbl">Data</div>
 <label for="dataset">HuggingFace dataset</label>
 <input type="text" id="dataset" value="HuggingFaceFW/fineweb-edu" spellcheck="false">
 <label for="subset">Config / subset <span style="font-weight:400">(blank = default)</span></label>
 <input type="text" id="subset" value="sample-10BT" spellcheck="false">
 <p class="sub" style="margin:.6rem 0 0;font-size:.82rem">Any streamable text dataset with a <code>text</code> column works. For gated/private datasets, log in first with <code>huggingface-cli login</code> on this machine — the trainer inherits your token.</p>
</div>
</div>
<div class="card" style="text-align:center">
 <button id="startbtn" onclick="start()">Start training</button>
 <button id="backbtn" onclick="goBack()" style="display:none;background:linear-gradient(135deg,#6b4423,#4a2f18)">&#8592; Back to settings</button>
 <p class="sub" style="margin:.6rem 0 0" id="status">idle</p></div>
<div class="card"><div class="lbl">Live loss</div><div class="num" id="loss">—</div></div>
<div class="card"><div class="lbl">Log</div><pre id="log"></pre></div>
<script>
const ids=["hidden","layers","heads","experts","seqlen","lr","batch","steps"];
ids.forEach(k=>{const el=document.getElementById(k);const v=document.getElementById("v"+k);
  el.oninput=()=>v.textContent=el.value;});
function cfg(){const c={};ids.forEach(k=>c[k]=+document.getElementById(k).value);c.lr=c.lr/1e4;
  c.dataset=document.getElementById("dataset").value.trim();
  c.subset=document.getElementById("subset").value.trim();return c;}
function showSettings(on){document.getElementById("settings").style.display=on?"":"none";
  document.getElementById("startbtn").style.display=on?"":"none";
  document.getElementById("backbtn").style.display=on?"none":"";}
async function start(){document.getElementById("status").textContent="starting…";
  showSettings(false);
  await fetch("/start",{method:"POST",body:JSON.stringify(cfg())});}
async function goBack(){await fetch("/stop",{method:"POST"});
  showSettings(true);document.getElementById("status").textContent="stopped — adjust and start again";}
async function poll(){try{const r=await fetch("/log");const d=await r.json();
  document.getElementById("log").textContent=d.log.slice().reverse().join("\\n");
  let last="—";for(const l of d.log){const m=l.match(/cluster-avg loss ([0-9.]+)/);if(m)last=m[1];}
  document.getElementById("loss").textContent=last;
  if(document.getElementById("backbtn").style.display!=="none")
    document.getElementById("status").textContent=d.running?"training…":"idle / done";}catch(e){}}
setInterval(poll,1000);poll();
</script></body></html>"""


class H(BaseHTTPRequestHandler):
    def _send(self, body, ctype="text/html; charset=utf-8", code=200):
        b = body.encode() if isinstance(body, str) else body
        self.send_response(code); self.send_header("Content-Type", ctype)
        self.send_header("Content-Length", str(len(b))); self.end_headers(); self.wfile.write(b)

    def do_GET(self):
        if self.path.startswith("/log"):
            with _lock:
                running = _proc is not None and _proc.poll() is None
                self._send(json.dumps({"log": list(_log), "running": running}), "application/json")
        else:
            self._send(PAGE)

    def do_POST(self):
        if self.path.startswith("/start"):
            n = int(self.headers.get("Content-Length", 0))
            cfg = json.loads(self.rfile.read(n) or "{}")
            ok, msg = start_training(cfg)
            self._send(json.dumps({"ok": ok, "msg": msg}), "application/json")
        elif self.path.startswith("/stop"):
            global _proc
            if _proc and _proc.poll() is None:
                _proc.terminate()
                with _lock:
                    _log.append("training stopped from the panel")
            self._send(json.dumps({"ok": True}), "application/json")

    def log_message(self, *a):
        pass


def main():
    print(f"[spikewhale-panel] http://localhost:{PORT}", flush=True)
    ThreadingHTTPServer(("0.0.0.0", PORT), H).serve_forever()


if __name__ == "__main__":
    main()