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:
@ -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()
|
||||
|
||||
Reference in New Issue
Block a user