diff --git a/model.onnx b/model.onnx index 4d86be3a..74081443 100644 Binary files a/model.onnx and b/model.onnx differ diff --git a/model.onnx.data b/model.onnx.data new file mode 100644 index 00000000..9beda747 Binary files /dev/null and b/model.onnx.data differ diff --git a/scripts/training_orchestrator.py b/scripts/training_orchestrator.py new file mode 100644 index 00000000..46c3a3a7 --- /dev/null +++ b/scripts/training_orchestrator.py @@ -0,0 +1,207 @@ +#!/usr/bin/env python3 +""" +持续训练 — 10k首轮 + 5k后续轮 → 循环到08:00 +每轮: dotnet自对弈 → V训练 → ONNX导出 +""" +import subprocess, os, sys, time, shutil, glob, json +from datetime import datetime + +HJHA_DIR = '/home/xiaoou/projects/hjha-server' +TRAINER_DIR = '/home/xiaoou/projects/paodekuai-trainer' +DEADLINE = '08:00' +LOG = os.path.join(HJHA_DIR, 'training_loop.log') + +def now(): return datetime.now().strftime('%H:%M:%S') +def log(msg): + line = f"[{now()}] {msg}" + print(line, flush=True) + with open(LOG, 'a') as f: f.write(line + '\n') + +def mins_left(): + dt = datetime.now() + h, m = map(int, DEADLINE.split(':')) + dl = dt.replace(hour=h, minute=m, second=0, microsecond=0) + return max(0, (dl - dt).total_seconds() / 60) + +def run(cmd, cwd=HJHA_DIR, to=14400): + """返回 (ok, stdout+stderr)""" + try: + r = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True, timeout=to) + ok = r.returncode == 0 + if not ok: + log(f" FAIL rc={r.returncode}: {(r.stderr or r.stdout)[:200]}") + return ok, r.stdout + r.stderr + except Exception as e: + log(f" EXCEPTION: {e}") + return False, str(e) + +# ── 自对弈 ── +def selfplay(games): + log(f"── 自对弈 {games}局 ──") + td = os.path.join(HJHA_DIR, 'training_data') + if os.path.exists(td): shutil.rmtree(td) + os.makedirs(td) + + env = os.environ.copy() + env['PATH'] = f"/usr/lib/dotnet:{env.get('PATH','')}" + + ok, out = run(['/usr/lib/dotnet/dotnet', 'run', '--project', 'hjha-console', + '--', '--mix', str(games)], to=max(14400, games*2)) + + csv = len(glob.glob(os.path.join(td, 'game_*.csv'))) + sf = os.path.join(HJHA_DIR, 'stress_summary.txt') + if os.path.exists(sf): + with open(sf) as f: log(f" {f.read().strip()}") + log(f" {csv} CSV") + return csv + +# ── V训练 ── +def train_v(): + log(f"── V训练 ──") + + sys.path.insert(0, TRAINER_DIR) + import numpy as np, torch + from models.network import PdkNet + + # load data + csv_dir = os.path.join(HJHA_DIR, 'training_data') + states, results = [], [] + for fp in sorted(glob.glob(f'{csv_dir}/game_*.csv')): + try: + with open(fp) as f: + f.readline() + for line in f: + p = line.strip().split('"') + if len(p) < 3: continue + sf = [float(x) for x in p[1].split(',')] + rest = p[2].strip(',').split(',') + if len(rest) < 2 or len(sf) != 364: continue + states.append(sf) + results.append(float(rest[1])) + except: pass + + if not states: + log(" 0 samples!") + return False + + states = np.array(states, dtype=np.float32) + results = np.array(results, dtype=np.float32) + log(f" {len(states)} samples from {len(glob.glob(f'{csv_dir}/game_*.csv'))} games") + + # load or init model — NOTE: old model.pt may be Q-network (14-output) + # V-training needs 1-output. Always start fresh for V-training. + mp = os.path.join(TRAINER_DIR, 'data/model.pt') + model = PdkNet() + log(" 新V模型 (从头训练)") + + opt = torch.optim.Adam(model.parameters(), lr=0.001) + loss_fn = torch.nn.MSELoss() + n = len(states) + + for ep in range(20): + perm = np.random.permutation(n) + losses = [] + for i in range(0, n, 128): + idx = perm[i:i+128] + x = torch.from_numpy(states[idx]).float() + y = torch.from_numpy(results[idx]).float().unsqueeze(1) + model.train() + pred = model(x) + loss = loss_fn(pred, y) + opt.zero_grad() + loss.backward() + opt.step() + losses.append(loss.item()) + if ep % 5 == 0: + log(f" ep{ep+1}/20 loss={np.mean(losses):.4f}") + log(f" done. final loss={np.mean(losses):.4f}") + + # save pytorch + model.save(mp) + + # export ONNX + try: + model.eval() + dummy = torch.randn(1, 364) + onnx_path = os.path.join(TRAINER_DIR, 'data/model.onnx') + torch.onnx.export(model, dummy, onnx_path, + input_names=['state'], output_names=['v'], + dynamic_axes={'state': {0: 'batch'}, 'v': {0: 'batch'}}, + opset_version=15) + + # copy to hjha-server (include .data if external storage) + dest = os.path.join(HJHA_DIR, 'model.onnx') + shutil.copy(onnx_path, dest) + data_src = onnx_path + '.data' + if os.path.exists(data_src): + shutil.copy(data_src, dest + '.data') + + # archive + tag = datetime.now().strftime('%m%d_%H%M') + rd = os.path.join(HJHA_DIR, f'training_rounds/r{tag}') + os.makedirs(rd, exist_ok=True) + shutil.copy(onnx_path, os.path.join(rd, 'model.onnx')) + if os.path.exists(data_src): + shutil.copy(data_src, os.path.join(rd, 'model.onnx.data')) + shutil.copy(mp, os.path.join(rd, 'model.pt')) + if os.path.exists(csv_dir): + shutil.move(csv_dir, os.path.join(rd, 'training_data')) + os.makedirs(csv_dir, exist_ok=True) + + sz = os.path.getsize(dest) / 1024 + log(f" ONNX → {dest} ({sz:.0f}KB), 存档 → {rd}") + except Exception as e: + log(f" ONNX失败: {e}") + return False + + return True + +# ── 主循环 ── +def main(): + log("=" * 50) + log(f"持续训练 启动 | 截止 {DEADLINE}") + log("策略: 首轮10k → 后续5k/轮") + log("=" * 50) + + rn, total = 0, 0 + while True: + ml = mins_left() + if ml < 3: + log(f"只剩 {ml:.0f}min, 停止") + break + + rn += 1 + games = 10000 if rn == 1 else 5000 + est = games/60 + 15 # estimate minutes + + if est > ml - 3: + # shrink + games = max(500, int((ml - 18) * 60)) + if games < 500: + log(f"不够时间跑完整轮 ({ml:.0f}min), 停止") + break + log(f"缩减为 {games} 局 (剩{ml:.0f}min)") + + log(f"\n── 第{rn}轮 {games}局 (剩{ml:.0f}min) ──") + + csv = selfplay(games) + if csv == 0: + log("自对弈失败! 跳过训练, 继续下一轮") + continue + total += games + + if not train_v(): + log("训练失败! 继续下一轮") + continue + + log(f"\n{'='*40}") + log(f"结束: {rn}轮, {total}局") + mp = os.path.join(HJHA_DIR, 'model.onnx') + if os.path.exists(mp): + log(f"最终模型: {mp} ({os.path.getsize(mp)/1024:.0f}KB)") + rd = os.path.join(HJHA_DIR, 'training_rounds') + if os.path.exists(rd): + log(f"存档: {len(os.listdir(rd))} 轮 → {rd}") + +if __name__ == '__main__': + main() \ No newline at end of file