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:
@@ -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)
|
||||
Reference in New Issue
Block a user