py -m rl.export vyexportuje checkpoint do rl/weights/neural-bot.json (bit-exact float32, 1.3 MB) a rl/pure_net.py ho hra bez torch/numpy (stdlib forward pass, ~16 ms/tah). Testy parity: logity aj akcie sa zhoduju s torch, identicke trajektorie celych kol. Natrenovany model: 6.7-6.9 b/kolo proti vsetkym baseline-om (heuristika prekonana). Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
156 lines
6.2 KiB
Python
156 lines
6.2 KiB
Python
"""Testy presnosti cisto-Python inferencie (rl/pure_net.py) voci torch.
|
|
|
|
Jadro suity: na zivych observaciach z nahodne rozohranych kol sa porovnavaju
|
|
logity a zvolene akcie pure-Python siete s torch sietou nacitanou z toho
|
|
isteho checkpointu. Case bez torch (cisty beh, legalnost, determinizmus)
|
|
bezia vzdy; porovnavacie case sa preskocia, ak torch nie je nainstalovany.
|
|
"""
|
|
|
|
import copy
|
|
import os
|
|
import unittest
|
|
from random import Random
|
|
|
|
from bridzik import Round
|
|
from rl.encoding import encode_observation, guess_mask, index_card, play_mask
|
|
from rl.env import PHASE_GUESS, RoundEnv
|
|
from rl.evaluate import play_round
|
|
from rl.players import RandomPlayer
|
|
from rl.pure_net import DEFAULT_WEIGHTS_PATH, PureNet, PureNeuralPlayer
|
|
|
|
WEIGHTS_AVAILABLE = os.path.exists(DEFAULT_WEIGHTS_PATH)
|
|
|
|
try:
|
|
import torch
|
|
from rl.policy_player import NeuralPlayer
|
|
from rl.train import load_checkpoint
|
|
TORCH_AVAILABLE = True
|
|
except ImportError: # pragma: no cover
|
|
TORCH_AVAILABLE = False
|
|
|
|
CHECKPOINT = os.path.join('rl', 'checkpoints', 'latest.pt')
|
|
|
|
|
|
def _random_decision_points(rng, n_rounds=12):
|
|
"""Vygeneruje zive rozhodovacie body (rnd, seat, faza) nahodnou hrou."""
|
|
env = RoundEnv(rng)
|
|
points = []
|
|
for i in range(n_rounds):
|
|
decision = env.reset(round_number=i % 8)
|
|
while True:
|
|
# snapshot -- env.round sa dalsou hrou mutuje
|
|
points.append((copy.deepcopy(env.round), decision.player, decision.phase))
|
|
action = rng.choice([a for a, ok in enumerate(decision.mask) if ok])
|
|
decision, rewards, done = env.step(action)
|
|
if done:
|
|
break
|
|
return points
|
|
|
|
|
|
@unittest.skipUnless(WEIGHTS_AVAILABLE, 'chyba export vah (py -m rl.export)')
|
|
class PureOnlyCase(unittest.TestCase):
|
|
"""Bezi aj bez torch -- presne to, co pobezi v produkcii."""
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.player = PureNeuralPlayer.load()
|
|
|
|
def test_plays_legal_full_rounds(self):
|
|
env = RoundEnv(Random(1))
|
|
players = [self.player, self.player,
|
|
RandomPlayer(Random(2)), RandomPlayer(Random(3))]
|
|
for round_number in range(8):
|
|
rewards = play_round(players, env, round_number)
|
|
self.assertEqual(len(rewards), 4)
|
|
|
|
def test_deterministic(self):
|
|
r = Round(2, 0)
|
|
self.assertEqual(self.player.guess(r, 0), self.player.guess(r, 0))
|
|
|
|
def test_respects_masks(self):
|
|
rng = Random(4)
|
|
for rnd, seat, phase in _random_decision_points(rng, n_rounds=8):
|
|
if phase == PHASE_GUESS:
|
|
self.assertTrue(guess_mask(rnd)[self.player.guess(rnd, seat)])
|
|
else:
|
|
self.assertTrue(play_mask(rnd, seat)[self.player.play(rnd, seat)])
|
|
|
|
|
|
@unittest.skipUnless(WEIGHTS_AVAILABLE and TORCH_AVAILABLE
|
|
and os.path.exists(CHECKPOINT),
|
|
'treba torch + checkpoint + export vah')
|
|
class TorchParityCase(unittest.TestCase):
|
|
"""Zhoda pure-Python inferencie s torch na tom istom checkpointe."""
|
|
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.pure = PureNet.load()
|
|
cls.torch_net = load_checkpoint(CHECKPOINT)
|
|
cls.torch_net.eval()
|
|
cls.points = _random_decision_points(Random(7), n_rounds=16)
|
|
|
|
def _torch_logits(self, obs, phase):
|
|
with torch.no_grad():
|
|
guess_logits, play_logits, _ = self.torch_net(
|
|
torch.tensor(obs, dtype=torch.float32).unsqueeze(0)
|
|
)
|
|
t = guess_logits if phase == PHASE_GUESS else play_logits
|
|
return t.squeeze(0).tolist()
|
|
|
|
def test_logits_match(self):
|
|
"""Logity sa zhoduju na ~1e-4 (rozdiel = len poradie scitovania
|
|
float32 vs float64, ziadna strata z exportu -- vahy su bit-exact)."""
|
|
worst = 0.0
|
|
for rnd, seat, phase in self.points:
|
|
obs = encode_observation(rnd, seat)
|
|
pure = self.pure.guess_logits(obs) if phase == PHASE_GUESS \
|
|
else self.pure.play_logits(obs)
|
|
ref = self._torch_logits(obs, phase)
|
|
for a, b in zip(pure, ref):
|
|
worst = max(worst, abs(a - b))
|
|
self.assertLess(worst, 1e-3, f'najvacsi rozdiel logitov: {worst}')
|
|
|
|
def test_actions_match(self):
|
|
"""Zvolena akcia je identicka vzdy, ked nejde o numericku remizu
|
|
(top-2 logity blizsie nez 1e-3 -- prakticky nenastava)."""
|
|
player = PureNeuralPlayer(self.pure)
|
|
torch_player = NeuralPlayer(self.torch_net, greedy=True)
|
|
compared = ties = 0
|
|
for rnd, seat, phase in self.points:
|
|
obs = encode_observation(rnd, seat)
|
|
if phase == PHASE_GUESS:
|
|
a, b = player.guess(rnd, seat), torch_player.guess(rnd, seat)
|
|
mask = guess_mask(rnd)
|
|
logits = self.pure.guess_logits(obs)
|
|
else:
|
|
a, b = player.play(rnd, seat), torch_player.play(rnd, seat)
|
|
mask = play_mask(rnd, seat)
|
|
logits = self.pure.play_logits(obs)
|
|
allowed = sorted((logits[i] for i in range(len(mask)) if mask[i]),
|
|
reverse=True)
|
|
if len(allowed) > 1 and allowed[0] - allowed[1] < 1e-3:
|
|
ties += 1 # numericka remiza -- volba je legitimne lubovolna
|
|
continue
|
|
compared += 1
|
|
self.assertEqual(a, b, f'akcie sa lisia mimo remizy ({phase})')
|
|
self.assertGreater(compared, 50) # test realne porovnaval
|
|
|
|
def test_full_rounds_identical_trajectories(self):
|
|
"""Dve identicke partie: pure aj torch hrac na vsetkych 4 sedadlach
|
|
s rovnakym rozdanim musia zahrat uplne rovnake kolo."""
|
|
pure_player = PureNeuralPlayer(self.pure)
|
|
torch_player = NeuralPlayer(self.torch_net, greedy=True)
|
|
for round_number in range(8):
|
|
results = []
|
|
for player in (pure_player, torch_player):
|
|
env = RoundEnv(Random(100 + round_number))
|
|
rewards = play_round([player] * 4, env, round_number)
|
|
results.append((rewards,
|
|
sorted(str(s.get_cards())
|
|
for s in env.round.stashes)))
|
|
self.assertEqual(results[0], results[1])
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main(verbosity=2)
|