feat: improve V-training loop — weight decay, cosine LR, early stop

- Adam weight_decay=1e-5 to prevent overfitting
- CosineAnnealingLR 0.001→1e-5 over 50 epochs
- Early stopping after 10 epochs without improvement
- Epochs: 20→50 max (with early stop typically fewer)
This commit is contained in:
2026-07-13 11:24:42 +08:00
parent 4c1b1c971c
commit 98ab4414a4

View File

@ -29,7 +29,7 @@ def run(cmd, cwd=HJHA_DIR, to=86400):
return False, str(e)
def selfplay(games):
log(f"── 自{games}局 ──")
log(f"── 自{games}局 ──")
td = os.path.join(HJHA_DIR, 'training_data')
if os.path.exists(td): shutil.rmtree(td)
os.makedirs(td)
@ -80,12 +80,15 @@ def train_v():
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)
opt = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-5)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=50, eta_min=1e-5)
loss_fn = torch.nn.MSELoss()
n = len(states)
t0 = datetime.now()
best_loss = float('inf')
no_improve = 0
for ep in range(20):
for ep in range(50):
perm = np.random.permutation(n)
losses = []
for i in range(0, n, 128):
@ -96,12 +99,23 @@ def train_v():
loss = loss_fn(pred, y)
opt.zero_grad(); loss.backward(); opt.step()
losses.append(loss.item())
if ep % 5 == 0:
scheduler.step()
avg_loss = np.mean(losses)
if avg_loss < best_loss * 0.999:
best_loss = avg_loss
no_improve = 0
else:
no_improve += 1
if ep % 5 == 0 or ep == 49:
dt = (datetime.now() - t0).total_seconds()
log(f" ep{ep+1}/20 loss={np.mean(losses):.4f} ({dt:.0f}s)")
lr = scheduler.get_last_lr()[0]
log(f" ep{ep+1}/50 loss={avg_loss:.4f} lr={lr:.6f} ({dt:.0f}s)")
if no_improve >= 10:
log(f" early stop ep{ep+1}, plateau 10 epochs")
break
elapsed = (datetime.now() - t0).total_seconds()
log(f" done {elapsed:.0f}s, final loss={np.mean(losses):.4f}")
log(f" done {elapsed:.0f}s, final loss={avg_loss:.4f}")
# save pt
mp = os.path.join(TRAINER_DIR, 'data/model.pt')
@ -155,7 +169,7 @@ def main():
csv = selfplay(games)
if csv == 0:
log("弈失败! 跳过")
log("弈失败! 跳过")
continue
total += games
@ -166,4 +180,4 @@ def main():
log(f"累计: {rn}轮, {total}")
if __name__ == '__main__':
main()
main()