Files
hjha-server/scripts/training_orchestrator.py

207 lines
6.8 KiB
Python

#!/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()