"""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()