Files
bridzik/tests/test_rl_train.py
T
timandClaude Fable 5 8f2449a408 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>
2026-07-07 18:49:55 +02:00

201 lines
8.0 KiB
Python

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