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)">← 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()
|