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