RL: self-play PPO trening siete

BridzikNet (trup + guess/play/value hlavy s maskovanim), self-play
generator so zdielanou sietou na 4 sedadlach a opponent mixingom
(random/heuristicke sedadla pre robustnost), vlastny clipped-PPO
so skalovanim odmien a lr/entropy annealom. Torch je len trenovacia
zavislost na hoste (requirements-rl.txt); checkpointy a logy su
gitignorovane. Spustenie: py -m rl.train.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
tim
2026-07-07 18:49:55 +02:00
co-authored by Claude Fable 5
parent e1733f4943
commit 8f2449a408
7 changed files with 659 additions and 0 deletions
+230
View File
@@ -0,0 +1,230 @@
"""Self-play PPO trening (viz rl/DESIGN.md, sekcie 4-6).
Slucka: nazbieraj self-play epizody -> PPO update -> kazdych par iteracii
evaluacia GREEDY policy proti fixnym baseline-om (nahodny hrac, MC heuristika)
z rl/players.py -- self-play reward sam o sebe nie je smerodajny (hra nie je
zero-sum, protihrac sa hybe spolu so sietou).
Spustenie:
py -m rl.train --iterations 200 --episodes 512
py -m rl.train --resume rl/checkpoints/latest.pt # pokracovanie
Checkpointy: rl/checkpoints/latest.pt (kazdu iteraciu) + best.pt (najlepsi
priemer bodov proti heuristikam). Metriky sa pripisuju do rl/runs/train_log.csv.
"""
import argparse
import csv
import os
import time
from random import Random
import torch
from rl.evaluate import evaluate
from rl.model import BridzikNet, masked_categorical
from rl.players import HeuristicPlayer, RandomPlayer
from rl.policy_player import NeuralPlayer
from rl.selfplay import collect_episodes
CHECKPOINT_DIR = os.path.join('rl', 'checkpoints')
RUNS_DIR = os.path.join('rl', 'runs')
def ppo_update(net: BridzikNet, optimizer: torch.optim.Optimizer, batch: dict,
clip: float = 0.2, epochs: int = 4, minibatch: int = 1024,
vf_coef: float = 1.0, ent_coef: float = 0.01,
max_grad_norm: float = 0.5) -> dict:
"""Standardny clipped-PPO krok nad batchom zo self-play.
Advantage sa standardizuje per batch (bod 7 v DESIGN.md -- odmeny 10-18 sa
lisia medzi kolami a zvysovali by varianciu gradientu). Guess a play kroky
zdielaju trup aj value hlavu, policy loss ide vzdy cez hlavu svojej fazy.
"""
n = batch['obs'].shape[0]
adv = batch['ret'] - batch['value']
adv = (adv - adv.mean()) / (adv.std() + 1e-8)
net.train()
stats = {'policy_loss': 0.0, 'value_loss': 0.0, 'entropy': 0.0, 'updates': 0}
for _ in range(epochs):
perm = torch.randperm(n)
for start in range(0, n, minibatch):
idx = perm[start:start + minibatch]
obs = batch['obs'][idx]
guess_logits, play_logits, value = net(obs)
is_play = batch['phase_play'][idx]
logp_new = torch.empty_like(batch['logp'][idx])
entropy = torch.empty_like(logp_new)
for phase_sel, logits, mask_key in (
(~is_play, guess_logits, 'guess_mask'),
(is_play, play_logits, 'play_mask'),
):
if not bool(phase_sel.any()):
continue
dist = masked_categorical(
logits[phase_sel], batch[mask_key][idx][phase_sel]
)
logp_new[phase_sel] = dist.log_prob(batch['action'][idx][phase_sel])
entropy[phase_sel] = dist.entropy()
ratio = torch.exp(logp_new - batch['logp'][idx])
mb_adv = adv[idx]
policy_loss = -torch.min(
ratio * mb_adv,
torch.clamp(ratio, 1 - clip, 1 + clip) * mb_adv,
).mean()
value_loss = (value - batch['ret'][idx]).pow(2).mean()
loss = policy_loss + vf_coef * value_loss - ent_coef * entropy.mean()
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(net.parameters(), max_grad_norm)
optimizer.step()
stats['policy_loss'] += policy_loss.item()
stats['value_loss'] += value_loss.item()
stats['entropy'] += entropy.mean().item()
stats['updates'] += 1
for key in ('policy_loss', 'value_loss', 'entropy'):
stats[key] /= max(stats['updates'], 1)
return stats
def evaluate_against_baselines(net: BridzikNet, n_rounds: int, rng: Random,
mc_samples: int = 60) -> dict:
"""Greedy siet na sedadle 0 vs 3x random a vs 3x heuristika."""
neural = NeuralPlayer(net, greedy=True)
vs_random = evaluate(
[neural] + [RandomPlayer(rng) for _ in range(3)], n_rounds, rng
)[0]
vs_heuristic = evaluate(
[neural] + [HeuristicPlayer(rng, n_samples=mc_samples) for _ in range(3)],
n_rounds, rng,
)[0]
return {
'vs_random_points': vs_random['avg_points'],
'vs_random_hit': vs_random['hit_rate'],
'vs_heuristic_points': vs_heuristic['avg_points'],
'vs_heuristic_hit': vs_heuristic['hit_rate'],
}
def save_checkpoint(net: BridzikNet, hidden: int, path: str) -> None:
torch.save({'hidden': hidden, 'state_dict': net.state_dict()}, path)
def load_checkpoint(path: str) -> BridzikNet:
"""Nacita checkpoint; podporuje aj stary format (bare state_dict)."""
payload = torch.load(path, map_location='cpu')
if isinstance(payload, dict) and 'state_dict' in payload:
net = BridzikNet(hidden=payload['hidden'])
net.load_state_dict(payload['state_dict'])
else:
net = BridzikNet()
net.load_state_dict(payload)
return net
def train(iterations: int, episodes: int, lr: float, seed: int,
eval_every: int, eval_rounds: int, resume: str = None,
hidden: int = 384, ent_coef_start: float = 0.01,
ent_coef_final: float = 0.001, lr_final_frac: float = 0.1,
mix_random: float = 0.0, mix_heuristic: float = 0.0):
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
os.makedirs(RUNS_DIR, exist_ok=True)
log_path = os.path.join(RUNS_DIR, 'train_log.csv')
log_exists = os.path.exists(log_path)
torch.manual_seed(seed)
rng = Random(seed)
if resume:
net = load_checkpoint(resume)
hidden = net.trunk[0].out_features
print(f'Pokracujem z checkpointu {resume} (hidden={hidden})')
else:
net = BridzikNet(hidden=hidden)
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
best_vs_heuristic = float('-inf')
with open(log_path, 'a', newline='') as log_file:
log = csv.writer(log_file)
if not log_exists:
log.writerow(['iteration', 'selfplay_points', 'policy_loss',
'value_loss', 'entropy', 'vs_random_points',
'vs_random_hit', 'vs_heuristic_points',
'vs_heuristic_hit', 'seconds'])
for iteration in range(1, iterations + 1):
started = time.time()
# linearny decay: lr klesa k lr*lr_final_frac, entropny bonus
# k ent_coef_final -- policy sa ku koncu behu moze doostrit
frac = 1 - (iteration - 1) / max(iterations - 1, 1)
for group in optimizer.param_groups:
group['lr'] = lr * (lr_final_frac + (1 - lr_final_frac) * frac)
ent_coef = ent_coef_final + (ent_coef_start - ent_coef_final) * frac
batch = collect_episodes(net, episodes, rng,
mix_random=mix_random,
mix_heuristic=mix_heuristic)
stats = ppo_update(net, optimizer, batch, ent_coef=ent_coef)
save_checkpoint(net, hidden, os.path.join(CHECKPOINT_DIR, 'latest.pt'))
row = [iteration, f'{batch["mean_points"]:.3f}',
f'{stats["policy_loss"]:.4f}', f'{stats["value_loss"]:.2f}',
f'{stats["entropy"]:.3f}']
line = (f'it {iteration:4d} | self-play {batch["mean_points"]:5.2f} '
f'b/kolo | pi {stats["policy_loss"]:+.4f} '
f'| V {stats["value_loss"]:7.2f} | H {stats["entropy"]:.3f}')
if iteration % eval_every == 0 or iteration == iterations:
ev = evaluate_against_baselines(net, eval_rounds, rng)
row += [f'{ev["vs_random_points"]:.3f}', f'{ev["vs_random_hit"]:.3f}',
f'{ev["vs_heuristic_points"]:.3f}', f'{ev["vs_heuristic_hit"]:.3f}']
line += (f' | vs random {ev["vs_random_points"]:5.2f} '
f'({100 * ev["vs_random_hit"]:.0f} %)'
f' | vs heur {ev["vs_heuristic_points"]:5.2f} '
f'({100 * ev["vs_heuristic_hit"]:.0f} %)')
if ev['vs_heuristic_points'] > best_vs_heuristic:
best_vs_heuristic = ev['vs_heuristic_points']
save_checkpoint(net, hidden,
os.path.join(CHECKPOINT_DIR, 'best.pt'))
line += ' *best*'
else:
row += ['', '', '', '']
row.append(f'{time.time() - started:.1f}')
log.writerow(row)
log_file.flush()
print(line)
return net
def main():
parser = argparse.ArgumentParser(description='Self-play PPO trening bridzik siete')
parser.add_argument('--iterations', type=int, default=200)
parser.add_argument('--episodes', type=int, default=512,
help='self-play kol na iteraciu')
parser.add_argument('--lr', type=float, default=3e-4)
parser.add_argument('--seed', type=int, default=1)
parser.add_argument('--eval-every', type=int, default=10)
parser.add_argument('--eval-rounds', type=int, default=400)
parser.add_argument('--hidden', type=int, default=384,
help='sirka skrytych vrstiev trupu')
parser.add_argument('--resume', type=str, default=None,
help='cesta k checkpointu (.pt) na pokracovanie')
parser.add_argument('--mix-random', type=float, default=0.0,
help='pravdepodobnost RandomPlayer sedadla v epizode')
parser.add_argument('--mix-heuristic', type=float, default=0.0,
help='pravdepodobnost HeuristicPlayer sedadla v epizode')
args = parser.parse_args()
train(args.iterations, args.episodes, args.lr, args.seed,
args.eval_every, args.eval_rounds, args.resume, hidden=args.hidden,
mix_random=args.mix_random, mix_heuristic=args.mix_heuristic)
if __name__ == '__main__':
main()