#!/usr/bin/env python3 """ 持续训练: 20k局/轮, 无限循环 每轮: dotnet自对弈 → V训练 → ONNX导出 → 存档 胜率趋势由 scan_wins.py 每小时 cron 独立采集 """ import subprocess, os, sys, time, shutil, glob from datetime import datetime HJHA_DIR = '/home/xiaoou/projects/hjha-server' TRAINER_DIR = '/home/xiaoou/projects/paodekuai-trainer' LOG = os.path.join(HJHA_DIR, 'training_loop.log') def now(): return datetime.now().strftime('%m%d %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 run(cmd, cwd=HJHA_DIR, to=86400): 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', '-c', 'Release', '--', '--mix', str(games)], to=max(36000, 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 def train_v(): log(f"── V训练 ──") sys.path.insert(0, TRAINER_DIR) import numpy as np, torch from models.network import PdkNet 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: if line.startswith('#'): continue 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") model = PdkNet() opt = torch.optim.Adam(model.parameters(), lr=0.001) loss_fn = torch.nn.MSELoss() n = len(states) t0 = datetime.now() 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: dt = (datetime.now() - t0).total_seconds() log(f" ep{ep+1}/20 loss={np.mean(losses):.4f} ({dt:.0f}s)") elapsed = (datetime.now() - t0).total_seconds() log(f" done {elapsed:.0f}s, final loss={np.mean(losses):.4f}") # save pt mp = os.path.join(TRAINER_DIR, 'data/model.pt') 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) 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') 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) + os.path.getsize(dest + '.data') if os.path.exists(dest + '.data') else os.path.getsize(dest) log(f" ONNX → {dest} ({sz//1024}KB), 存档 → {rd}") except Exception as e: log(f" ONNX failed: {e}") return False return True def main(): log("=" * 50) log("持续训练: 20k局/轮 | 无限循环") log("=" * 50) rn, total = 0, 0 while True: rn += 1 games = 20000 log(f"\n── 第{rn}轮: {games}局 ──") csv = selfplay(games) if csv == 0: log("自对弈失败! 跳过") continue total += games if not train_v(): log("训练失败! 跳过") continue log(f"累计: {rn}轮, {total}局") if __name__ == '__main__': main()