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