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
+2
View File
@@ -10,3 +10,5 @@ frontend/.vite/
.env.*
!.env.example
geoip/*.mmdb
rl/runs/
rl/checkpoints/
+5
View File
@@ -0,0 +1,5 @@
# RL trening (rl/model.py, rl/selfplay.py, rl/train.py) -- zamerne oddelene
# od requirements.txt: server ani Docker image torch nepotrebuju, boti v hre
# pouzivaju len cisto-Python rl/players.py (a neskor natrenovane vahy cez
# torch az ked sa neuralny bot nasadi).
torch>=2.4
+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()
+200
View File
@@ -0,0 +1,200 @@
"""Testy PPO pipeline (rl/model, rl/selfplay, rl/policy_player, rl/train).
Vyzaduju torch (requirements-rl.txt); bez neho sa cely modul preskoci --
ostatne suity (engine, encoding, boti) na torchi nezavisia.
"""
import unittest
from random import Random
try:
import torch
except ImportError: # pragma: no cover
raise unittest.SkipTest('torch nie je nainstalovany (requirements-rl.txt)')
from bridzik import Round
from rl.encoding import (
N_GUESS_ACTIONS, N_PLAY_ACTIONS, OBS_DIM,
encode_observation, guess_mask, play_mask,
)
from rl.evaluate import evaluate, play_round
from rl.env import RoundEnv
from rl.model import BridzikNet, mask_tensor, masked_categorical, obs_tensor
from rl.players import RandomPlayer
from rl.policy_player import NeuralPlayer
from rl.selfplay import _assign_seats, collect_episodes
from rl.train import ppo_update
class ModelCase(unittest.TestCase):
def test_output_shapes(self):
net = BridzikNet(hidden=32)
obs = torch.zeros((5, OBS_DIM))
guess_logits, play_logits, value = net(obs)
self.assertEqual(guess_logits.shape, (5, N_GUESS_ACTIONS))
self.assertEqual(play_logits.shape, (5, N_PLAY_ACTIONS))
self.assertEqual(value.shape, (5,))
def test_masked_categorical_never_samples_illegal(self):
torch.manual_seed(0)
logits = torch.zeros((1, 9))
mask = torch.tensor([[False, True, False, True, False,
False, False, False, False]])
dist = masked_categorical(logits, mask)
samples = dist.sample((200,))
self.assertTrue(set(samples.flatten().tolist()) <= {1, 3})
# entropia a log_prob su konecne aj s -inf logitmi
self.assertTrue(torch.isfinite(dist.entropy()).all())
self.assertTrue(torch.isfinite(dist.log_prob(torch.tensor([1]))).all())
def test_encoding_tensors(self):
r = Round(3, 0)
obs = obs_tensor(encode_observation(r, 0))
self.assertEqual(obs.shape, (OBS_DIM,))
self.assertEqual(mask_tensor(guess_mask(r)).shape, (N_GUESS_ACTIONS,))
self.assertEqual(mask_tensor(play_mask(r, 0)).shape, (N_PLAY_ACTIONS,))
class NeuralPlayerCase(unittest.TestCase):
def test_untrained_net_plays_legal_full_rounds(self):
torch.manual_seed(1)
net = BridzikNet(hidden=32)
env = RoundEnv(Random(2))
players = [NeuralPlayer(net, greedy=True),
NeuralPlayer(net, greedy=False),
RandomPlayer(Random(3)), RandomPlayer(Random(4))]
# dohratie kola bez BridzikException = vsetky tahy legalne
for round_number in range(8):
rewards = play_round(players, env, round_number)
self.assertEqual(len(rewards), 4)
class SelfPlayCase(unittest.TestCase):
def test_collect_episodes_batch_consistency(self):
torch.manual_seed(5)
net = BridzikNet(hidden=32)
batch = collect_episodes(net, n_episodes=6, rng=Random(6))
n = batch['obs'].shape[0]
self.assertGreater(n, 0)
for key, width in (('guess_mask', N_GUESS_ACTIONS),
('play_mask', N_PLAY_ACTIONS)):
self.assertEqual(batch[key].shape, (n, width))
for key in ('phase_play', 'action', 'logp', 'value', 'ret'):
self.assertEqual(batch[key].shape, (n,))
# kazda epizoda ma prave 4 guess kroky -> pocet guess krokov = 4*epizody
self.assertEqual(int((~batch['phase_play']).sum()), 4 * 6)
# return je bud 0 alebo (10+tip)/REWARD_SCALE, cize v (0.55, 1.0]
for r in batch['ret'].tolist():
self.assertTrue(r == 0.0 or 10.0 / 18.0 <= r <= 1.0)
# akcia bola vzdy legalna podla ulozenej masky svojej fazy
for i in range(n):
mask = batch['play_mask'][i] if batch['phase_play'][i] \
else batch['guess_mask'][i]
self.assertTrue(bool(mask[batch['action'][i]]))
self.assertGreaterEqual(batch['mean_points'], 0.0)
class OpponentMixingCase(unittest.TestCase):
def test_assign_seats_always_keeps_a_net_seat(self):
rng = Random(20)
marker = object()
for _ in range(200):
seats = _assign_seats(rng, 1.0, 0.0, marker, marker)
self.assertIn(None, seats.values()) # aj pri mix_random=1.0
self.assertEqual(set(seats), {0, 1, 2, 3})
def test_mixed_episodes_record_only_net_seats(self):
torch.manual_seed(21)
net = BridzikNet(hidden=32)
# mix_random=1.0 -> presne jedno sietove sedadlo na epizodu
batch = collect_episodes(net, n_episodes=5, rng=Random(22),
mix_random=1.0)
# 1 sietove sedadlo = presne 1 guess krok na epizodu
self.assertEqual(int((~batch['phase_play']).sum()), 5)
for i in range(batch['obs'].shape[0]):
mask = batch['play_mask'][i] if batch['phase_play'][i] \
else batch['guess_mask'][i]
self.assertTrue(bool(mask[batch['action'][i]]))
for r in batch['ret'].tolist():
self.assertTrue(r == 0.0 or 10.0 / 18.0 <= r <= 1.0)
def test_mixed_episodes_with_heuristic(self):
torch.manual_seed(23)
net = BridzikNet(hidden=32)
batch = collect_episodes(net, n_episodes=4, rng=Random(24),
mix_heuristic=0.5, heuristic_samples=10)
n_guess = int((~batch['phase_play']).sum())
self.assertGreaterEqual(n_guess, 4) # aspon 1 sietove sedadlo/epizodu
self.assertLessEqual(n_guess, 16)
self.assertGreaterEqual(batch['mean_points'], 0.0)
class PpoUpdateCase(unittest.TestCase):
def test_update_changes_params_and_is_finite(self):
torch.manual_seed(7)
net = BridzikNet(hidden=32)
optimizer = torch.optim.Adam(net.parameters(), lr=1e-3)
batch = collect_episodes(net, n_episodes=8, rng=Random(8))
before = [p.detach().clone() for p in net.parameters()]
stats = ppo_update(net, optimizer, batch, epochs=2, minibatch=64)
for key in ('policy_loss', 'value_loss', 'entropy'):
self.assertTrue(torch.isfinite(torch.tensor(stats[key])))
changed = any(
not torch.equal(b, a.detach())
for b, a in zip(before, net.parameters())
)
self.assertTrue(changed)
def test_value_head_learns_constant_reward(self):
# sanity uciaceho kroku: na batchi s konstantnym returnom sa value
# loss po par updatoch zmensi
torch.manual_seed(9)
net = BridzikNet(hidden=32)
optimizer = torch.optim.Adam(net.parameters(), lr=3e-3)
batch = collect_episodes(net, n_episodes=8, rng=Random(10))
batch['ret'] = torch.full_like(batch['ret'], 12.0 / 18.0)
first = ppo_update(net, optimizer, batch, epochs=1, minibatch=4096)
for _ in range(10):
last = ppo_update(net, optimizer, batch, epochs=1, minibatch=4096)
self.assertLess(last['value_loss'], first['value_loss'])
class CheckpointCase(unittest.TestCase):
def test_save_load_roundtrip_with_hidden(self):
import os
import tempfile
from rl.train import load_checkpoint, save_checkpoint
torch.manual_seed(13)
net = BridzikNet(hidden=48)
path = os.path.join(tempfile.gettempdir(), 'bridzik_ckpt_test.pt')
try:
save_checkpoint(net, 48, path)
loaded = load_checkpoint(path)
self.assertEqual(loaded.trunk[0].out_features, 48)
obs = torch.zeros((1, OBS_DIM))
for a, b in zip(net(obs), loaded(obs)):
self.assertTrue(torch.equal(a, b))
finally:
os.remove(path)
class EvaluateIntegrationCase(unittest.TestCase):
def test_neural_player_in_harness(self):
torch.manual_seed(11)
net = BridzikNet(hidden=32)
rng = Random(12)
stats = evaluate(
[NeuralPlayer(net)] + [RandomPlayer(rng) for _ in range(3)],
n_rounds=20, rng=rng,
)
self.assertEqual(stats[0]['rounds'], 20)
if __name__ == '__main__':
unittest.main(verbosity=2)