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
+49
View File
@@ -0,0 +1,49 @@
"""Siet pre self-play PPO (viz rl/DESIGN.md, sekcia 3).
Zdielany trup nad observaciou z rl/encoding.py a tri hlavy:
guess (9 logitov), play (32 logitov), value (1 skalar -- baseline pre
actor-critic). Ktora policy hlava plati, urcuje faza rozhodnutia; nelegalne
akcie sa odrezavaju maskou (logit -inf), takze distribucia nikdy nenavzorkuje
tah, ktory by engine odmietol.
"""
import torch
import torch.nn as nn
from rl.encoding import N_GUESS_ACTIONS, N_PLAY_ACTIONS, OBS_DIM
class BridzikNet(nn.Module):
def __init__(self, hidden: int = 256):
super().__init__()
self.trunk = nn.Sequential(
nn.Linear(OBS_DIM, hidden), nn.ReLU(),
nn.Linear(hidden, hidden), nn.ReLU(),
)
self.guess_head = nn.Linear(hidden, N_GUESS_ACTIONS)
self.play_head = nn.Linear(hidden, N_PLAY_ACTIONS)
self.value_head = nn.Linear(hidden, 1)
def forward(self, obs: torch.Tensor):
"""obs (B, OBS_DIM) -> (guess_logits (B,9), play_logits (B,32), value (B,))."""
h = self.trunk(obs)
return self.guess_head(h), self.play_head(h), self.value_head(h).squeeze(-1)
def masked_categorical(logits: torch.Tensor, mask: torch.Tensor) -> torch.distributions.Categorical:
"""Kategoricka distribucia s nelegalnymi akciami odrezanymi na -inf.
`mask` je bool tensor rovnakeho tvaru ako `logits`; kazdy riadok musi mat
aspon jednu povolenu akciu (garantuju masky z rl/encoding.py).
"""
return torch.distributions.Categorical(
logits=logits.masked_fill(~mask, float('-inf'))
)
def obs_tensor(obs: list) -> torch.Tensor:
return torch.tensor(obs, dtype=torch.float32)
def mask_tensor(mask: list) -> torch.Tensor:
return torch.tensor(mask, dtype=torch.bool)
+42
View File
@@ -0,0 +1,42 @@
"""Natrenovana siet ako hrac so standardnym rozhranim guess/play.
Rovnake rozhranie ako rl/players.py, takze funguje v rl/evaluate.py aj ako
boti "mozog" v api/bots.py. Hrac vidi len observaciu + masku z rl/encoding.py
-- z principu nemoze podvadzat (do cudzich ruk sa nedostane).
"""
import torch
from rl.encoding import encode_observation, guess_mask, play_mask
from rl.model import BridzikNet, mask_tensor, masked_categorical, obs_tensor
class NeuralPlayer:
def __init__(self, net: BridzikNet, greedy: bool = True):
self.net = net
self.greedy = greedy # argmax pri evaluacii; sampling pre pestrost
def _act(self, rnd, seat: int, use_play_head: bool) -> int:
obs = obs_tensor(encode_observation(rnd, seat)).unsqueeze(0)
mask = mask_tensor(
play_mask(rnd, seat) if use_play_head else guess_mask(rnd)
).unsqueeze(0)
self.net.eval()
with torch.no_grad():
guess_logits, play_logits, _ = self.net(obs)
logits = play_logits if use_play_head else guess_logits
logits = logits.masked_fill(~mask, float('-inf'))
if self.greedy:
return int(logits.argmax(dim=-1).item())
return int(masked_categorical(logits, mask).sample().item())
def guess(self, rnd, seat: int) -> int:
return self._act(rnd, seat, use_play_head=False)
def play(self, rnd, seat: int) -> int:
return self._act(rnd, seat, use_play_head=True)
def load_player(checkpoint_path: str, greedy: bool = True) -> NeuralPlayer:
from rl.train import load_checkpoint
return NeuralPlayer(load_checkpoint(checkpoint_path), greedy=greedy)
+131
View File
@@ -0,0 +1,131 @@
"""Self-play generator: jedna zdielana siet hra vsetkych 4 hracov v Round
epizodach a zbiera trajektorie pre PPO (viz rl/DESIGN.md, sekcie 4-5).
Odmena je sparse a terminalna: kazde rozhodnutie hraca v kole (tip aj vsetky
karty) dostane ako return jeho `points_summary` z konca kola, gamma = 1.
Masky sa ukladaju oddelene pre obe fazy (rozne velkosti akcneho priestoru);
`phase_play` hovori, ktora hlava/maska pre dany krok plati.
"""
from random import Random
import torch
from rl.encoding import N_GUESS_ACTIONS, N_PLAY_ACTIONS
from rl.env import PHASE_PLAY, RoundEnv
from rl.model import BridzikNet, mask_tensor, masked_categorical, obs_tensor
from rl.players import HeuristicPlayer, RandomPlayer
# Returny sa skaluju do [0, 1] (max odmena je 10+8). Bez skalovania ma value
# loss (MSE na 0-18) radovo vacsi gradient nez policy loss a cez zdielany
# trup policy ucenie prevalcuje.
REWARD_SCALE = 18.0
def _assign_seats(rng: Random, mix_random: float, mix_heuristic: float,
random_player, heuristic_player) -> dict:
"""Obsadenie sedadiel pre jednu epizodu: None = siet, inak skriptovany
supper. Aspon jedno sedadlo musi hrat siet (inak niet co zbierat)."""
seats = {}
for seat in range(4):
roll = rng.random()
if roll < mix_random:
seats[seat] = random_player
elif roll < mix_random + mix_heuristic:
seats[seat] = heuristic_player
else:
seats[seat] = None
if not any(p is None for p in seats.values()):
seats[rng.randrange(4)] = None
return seats
def collect_episodes(net: BridzikNet, n_episodes: int, rng: Random,
round_numbers: list = None, mix_random: float = 0.0,
mix_heuristic: float = 0.0,
heuristic_samples: int = 40) -> dict:
"""Odohra `n_episodes` self-play kol a vrati batch tenzorov:
obs (N, OBS_DIM), phase_play (N,) bool, action (N,), logp (N,), value (N,),
ret (N,), guess_mask (N, 9), play_mask (N, 32) -- maska nepatriacej fazy je
pre dany krok cela False a pri update sa nepouzije.
Navyse 'mean_points': priemerne body na sietove sedadlo a kolo.
Opponent mixing (robustnost na nie-self-play superov): s pravdepodobnostou
`mix_random` / `mix_heuristic` hra sedadlo RandomPlayer / HeuristicPlayer
namiesto siete. Tahy skriptovanych superov sa do batchu NEZAZNAMENAVAJU
(nie su z trenovanej policy) -- superi len obsadzuju stol.
"""
env = RoundEnv(rng)
random_player = RandomPlayer(rng)
heuristic_player = HeuristicPlayer(rng, n_samples=heuristic_samples)
mixing = mix_random > 0 or mix_heuristic > 0
obs_l, phase_l, action_l, logp_l, value_l, ret_l = [], [], [], [], [], []
gmask_l, pmask_l = [], []
total_points = 0.0
net_seat_rounds = 0
net.eval()
with torch.no_grad():
for _ in range(n_episodes):
round_number = rng.choice(round_numbers) if round_numbers else None
decision = env.reset(round_number)
seats = _assign_seats(rng, mix_random, mix_heuristic,
random_player, heuristic_player) if mixing \
else {seat: None for seat in range(4)}
net_seat_rounds += sum(1 for p in seats.values() if p is None)
# indexy krokov sietovych sedadiel -- na priradenie returnu
player_steps = {p: [] for p in range(4) if seats[p] is None}
while True:
opponent = seats[decision.player]
if opponent is not None:
# skriptovany supper: vykonaj tah, nic nezaznamenavaj
if decision.phase == PHASE_PLAY:
action_i = opponent.play(env.round, decision.player)
else:
action_i = opponent.guess(env.round, decision.player)
else:
obs = obs_tensor(decision.obs).unsqueeze(0)
mask = mask_tensor(decision.mask).unsqueeze(0)
guess_logits, play_logits, value = net(obs)
is_play = decision.phase == PHASE_PLAY
dist = masked_categorical(
play_logits if is_play else guess_logits, mask
)
action = dist.sample()
action_i = action.item()
player_steps[decision.player].append(len(obs_l))
obs_l.append(decision.obs)
phase_l.append(is_play)
action_l.append(action_i)
logp_l.append(dist.log_prob(action).item())
value_l.append(value.item())
ret_l.append(0.0) # doplni sa na konci kola
if is_play:
gmask_l.append([False] * N_GUESS_ACTIONS)
pmask_l.append(decision.mask)
else:
gmask_l.append(decision.mask)
pmask_l.append([False] * N_PLAY_ACTIONS)
decision, rewards, done = env.step(action_i)
if done:
for player, steps in player_steps.items():
for i in steps:
ret_l[i] = rewards[player] / REWARD_SCALE
total_points += rewards[player]
break
return {
'obs': torch.tensor(obs_l, dtype=torch.float32),
'phase_play': torch.tensor(phase_l, dtype=torch.bool),
'action': torch.tensor(action_l, dtype=torch.long),
'logp': torch.tensor(logp_l, dtype=torch.float32),
'value': torch.tensor(value_l, dtype=torch.float32),
'ret': torch.tensor(ret_l, dtype=torch.float32),
'guess_mask': torch.tensor(gmask_l, dtype=torch.bool),
'play_mask': torch.tensor(pmask_l, dtype=torch.bool),
'mean_points': total_points / max(net_seat_rounds, 1),
}
+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()