diff --git a/.gitignore b/.gitignore index 4f24611..1e0e272 100644 --- a/.gitignore +++ b/.gitignore @@ -10,3 +10,5 @@ frontend/.vite/ .env.* !.env.example geoip/*.mmdb +rl/runs/ +rl/checkpoints/ diff --git a/requirements-rl.txt b/requirements-rl.txt new file mode 100644 index 0000000..8e74032 --- /dev/null +++ b/requirements-rl.txt @@ -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 diff --git a/rl/model.py b/rl/model.py new file mode 100644 index 0000000..c886195 --- /dev/null +++ b/rl/model.py @@ -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) diff --git a/rl/policy_player.py b/rl/policy_player.py new file mode 100644 index 0000000..f2f3a32 --- /dev/null +++ b/rl/policy_player.py @@ -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) diff --git a/rl/selfplay.py b/rl/selfplay.py new file mode 100644 index 0000000..cb66047 --- /dev/null +++ b/rl/selfplay.py @@ -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), + } diff --git a/rl/train.py b/rl/train.py new file mode 100644 index 0000000..3ddf6cb --- /dev/null +++ b/rl/train.py @@ -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() diff --git a/tests/test_rl_train.py b/tests/test_rl_train.py new file mode 100644 index 0000000..6a63faf --- /dev/null +++ b/tests/test_rl_train.py @@ -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)