169 lines
5.6 KiB
Python
169 lines
5.6 KiB
Python
#!/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() |