Files
bridzik/tests/test_encoding.py
T
timandClaude Fable 5 3710a68e37 RL: encoding observacii, akcne masky a Round prostredie
Zaklad RL vrstvy (rl/DESIGN.md): egocentricky rotovana observacia
(ruka, videne karty, tipy, kopka, dedukovane voidy -- 233 dim), masky
legalnych tipov/kariet zrkadliace pravidla enginu a RoundEnv (jedno kolo
= jedna self-play epizoda). Fuzz-testy vynucuju zhodu masiek s enginom.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-07 18:49:55 +02:00

297 lines
13 KiB
Python

import copy
import random
import unittest
from bridzik import cards, Card, Card_colors, Card_values, BridzikException, Round
from rl.encoding import (
COLORS, VALUES, N_CARDS, N_GUESS_ACTIONS, N_PLAY_ACTIONS,
OFF_HAND, OFF_SEEN, OFF_ROUND, OFF_PHASE, OFF_GUESSES, OFF_TRICKS,
OFF_STASH, OFF_STASH_LEADER, OFF_VOIDS, OBS_DIM,
card_index, deduce_voids, index_card, relative_seat, encode_observation,
guess_mask, play_mask,
)
class CardIndexCase(unittest.TestCase):
def test_roundtrip_and_uniqueness(self):
indexes = set()
for card in cards:
idx = card_index(card)
self.assertIn(idx, range(N_CARDS))
self.assertEqual(index_card(idx), card)
indexes.add(idx)
self.assertEqual(len(indexes), N_CARDS)
def test_layout(self):
# farba = blok po 8, hodnota = pozicia v bloku
self.assertEqual(card_index(Card(Card_colors['HEARTS'], Card_values['C7'])), 0)
self.assertEqual(card_index(Card(Card_colors['HEARTS'], Card_values['ACE'])), 7)
self.assertEqual(card_index(Card(COLORS[3], Card_values['ACE'])), 31)
class RotationCase(unittest.TestCase):
def test_relative_seat(self):
for player in range(4):
self.assertEqual(relative_seat(player, player), 0)
# smer hry = rastuce cislo sedadla mod 4
self.assertEqual(relative_seat((player + 1) % 4, player), 1)
self.assertEqual(relative_seat((player + 3) % 4, player), 3)
def test_guesses_rotated_for_all_seats(self):
r = Round(0, 2)
guesses = {2: 5, 3: 0, 0: 1, 1: 1}
for seat in [2, 3, 0, 1]:
r.add_player_guess(seat, guesses[seat])
for player in range(4):
obs = encode_observation(r, player)
for seat in range(4):
rel = relative_seat(seat, player)
self.assertEqual(obs[OFF_GUESSES + 2 * rel], 1.0)
self.assertEqual(obs[OFF_GUESSES + 2 * rel + 1], guesses[seat] / 8)
def test_partial_guesses_flags(self):
r = Round(3, 1)
r.add_player_guess(1, 2)
for player in range(4):
obs = encode_observation(r, player)
rel = relative_seat(1, player)
self.assertEqual(obs[OFF_GUESSES + 2 * rel], 1.0)
self.assertEqual(obs[OFF_GUESSES + 2 * rel + 1], 2 / 8)
for seat in [0, 2, 3]:
rel = relative_seat(seat, player)
self.assertEqual(obs[OFF_GUESSES + 2 * rel], 0.0)
self.assertEqual(obs[OFF_GUESSES + 2 * rel + 1], 0.0)
class ObservationCase(unittest.TestCase):
@staticmethod
def _deterministic_round():
# rovnaka konstrukcia ako v test_engine.RoundCase.test_play_card
shuffler = lambda l: None
c0 = [Card(Card_colors['BELLS'], Card_values['UPPER']),
Card(Card_colors['HEARTS'], Card_values['UPPER'])]
c1 = [Card(Card_colors['BELLS'], Card_values['C7']),
Card(Card_colors['HEARTS'], Card_values['C10'])]
c2 = [Card(Card_colors['BELLS'], Card_values['ACE']),
Card(Card_colors['BELLS'], Card_values['C8'])]
c3 = [Card(Card_colors['LEAVES'], Card_values['C7']),
Card(Card_colors['BELLS'], Card_values['LOWER'])]
c = ['dummy'] * 24 + c0 + c1 + c2 + c3
r = Round(6, 1, c, shuffler)
return r, [c0, c1, c2, c3]
def test_hand_multi_hot(self):
r, hands = self._deterministic_round()
for player in range(4):
obs = encode_observation(r, player)
hand_indexes = {card_index(c) for c in hands[player]}
for i in range(N_CARDS):
self.assertEqual(obs[OFF_HAND + i], 1.0 if i in hand_indexes else 0.0)
def test_round_number_and_phase(self):
r, _ = self._deterministic_round()
obs = encode_observation(r, 0)
for i in range(8):
self.assertEqual(obs[OFF_ROUND + i], 1.0 if i == 6 else 0.0)
self.assertEqual(obs[OFF_PHASE], 1.0) # tipovacia faza
for seat, guess in [(1, 0), (2, 0), (3, 1), (0, 2)]:
r.add_player_guess(seat, guess)
obs = encode_observation(r, 0)
self.assertEqual(obs[OFF_PHASE], 0.0) # hracia faza
def test_current_stash_slots_and_seen(self):
r, hands = self._deterministic_round()
for seat, guess in [(1, 0), (2, 0), (3, 1), (0, 2)]:
r.add_player_guess(seat, guess)
# rozohrana kopka: hraju 0 a 1
r.play_card(0, hands[0][0])
r.play_card(1, hands[1][0])
for player in range(4):
obs = encode_observation(r, player)
slot0 = relative_seat(0, player)
slot1 = relative_seat(1, player)
self.assertEqual(obs[OFF_STASH + slot0 * N_CARDS + card_index(hands[0][0])], 1.0)
self.assertEqual(obs[OFF_STASH + slot1 * N_CARDS + card_index(hands[1][0])], 1.0)
self.assertEqual(sum(obs[OFF_STASH:OFF_STASH + 4 * N_CARDS]), 2.0)
# leader kopky je hrac 0 (najvyssi tip)
self.assertEqual(obs[OFF_STASH_LEADER + relative_seat(0, player)], 1.0)
# nic este nie je "videne" -- prva kopka nie je dokoncena
self.assertEqual(sum(obs[OFF_SEEN:OFF_SEEN + N_CARDS]), 0.0)
# dokoncena kopka -> karty sa presunu do SEEN, sloty sa vyprazdnia
r.play_card(2, hands[2][0])
r.play_card(3, hands[3][1])
obs = encode_observation(r, 0)
played = [hands[0][0], hands[1][0], hands[2][0], hands[3][1]]
for card in played:
self.assertEqual(obs[OFF_SEEN + card_index(card)], 1.0)
self.assertEqual(sum(obs[OFF_SEEN:OFF_SEEN + N_CARDS]), 4.0)
self.assertEqual(sum(obs[OFF_STASH:OFF_STASH + 4 * N_CARDS]), 0.0)
# novu kopku vynasa vitaz (hrac 2, BELLS ACE)
self.assertEqual(obs[OFF_STASH_LEADER + relative_seat(2, 0)], 1.0)
# pocty vyhranych kopiek rotovane
for player in range(4):
obs = encode_observation(r, player)
self.assertEqual(obs[OFF_TRICKS + relative_seat(2, player)], 1 / 8)
def test_terminal_state_encodable(self):
r = Round(7, 0)
for seat, guess in [(0, 0), (1, 0), (2, 0), (3, 0)]:
try:
r.add_player_guess(seat, guess)
except BridzikException:
r.add_player_guess(seat, 1)
while not r.is_completed():
player = r.get_active_player()
mask = play_mask(r, player)
r.play_card(player, index_card(mask.index(True)))
obs = encode_observation(r, 0)
self.assertEqual(len(obs), OBS_DIM)
self.assertEqual(sum(obs[OFF_SEEN:OFF_SEEN + N_CARDS]), 4.0)
class VoidsInObservationCase(unittest.TestCase):
def test_voids_encoded_and_rotated(self):
# hrac 0 vynasa zelen; 2 tromfne (void zelen), 3 hodi gulu (void
# zelen aj cerven) -- viz deduce_voids
hand0 = [Card(Card_colors['LEAVES'], Card_values['C7']),
Card(Card_colors['LEAVES'], Card_values['C8'])]
hand1 = [Card(Card_colors['LEAVES'], Card_values['C9']),
Card(Card_colors['LEAVES'], Card_values['C10'])]
hand2 = [Card(Card_colors['HEARTS'], Card_values['C7']),
Card(Card_colors['ACORNS'], Card_values['C7'])]
hand3 = [Card(Card_colors['BELLS'], Card_values['C7']),
Card(Card_colors['BELLS'], Card_values['C8'])]
rest = [c for c in cards if c not in hand0 + hand1 + hand2 + hand3]
deck = rest[:24] + hand0 + hand1 + hand2 + hand3
r = Round(6, 0, deck, shuffler=lambda l: None)
r.add_player_guess(0, 2)
r.add_player_guess(1, 0)
r.add_player_guess(2, 0)
r.add_player_guess(3, 1)
obs = encode_observation(r, 0)
self.assertEqual(sum(obs[OFF_VOIDS:OFF_VOIDS + 16]), 0.0)
for seat, card in [(0, hand0[0]), (1, hand1[0]),
(2, hand2[0]), (3, hand3[0])]:
r.play_card(seat, card)
leaves_i = COLORS.index(Card_colors['LEAVES'])
hearts_i = COLORS.index(Card_colors['HEARTS'])
for player in range(4):
obs = encode_observation(r, player)
block = lambda seat: obs[OFF_VOIDS + relative_seat(seat, player) * 4:
OFF_VOIDS + relative_seat(seat, player) * 4 + 4]
self.assertEqual(sum(block(0)), 0.0) # vynasajuci neprezradza nic
self.assertEqual(sum(block(1)), 0.0) # priznal farbu
self.assertEqual(block(2)[leaves_i], 1.0)
self.assertEqual(sum(block(2)), 1.0)
self.assertEqual(block(3)[leaves_i], 1.0)
self.assertEqual(block(3)[hearts_i], 1.0)
self.assertEqual(sum(block(3)), 2.0)
class GuessMaskCase(unittest.TestCase):
def test_range_by_round_number(self):
for round_number in range(8):
r = Round(round_number, 0)
mask = guess_mask(r)
for g in range(N_GUESS_ACTIONS):
self.assertEqual(mask[g], g <= 8 - round_number)
def test_last_guesser_forbidden_value(self):
r = Round(0, 0)
r.add_player_guess(0, 2)
r.add_player_guess(1, 1)
r.add_player_guess(2, 3)
mask = guess_mask(r)
self.assertFalse(mask[2]) # 2+1+3+2 == 8 kopiek -> zakazane
for g in [0, 1, 3, 4, 5, 6, 7, 8]:
self.assertTrue(mask[g])
def test_forbidden_value_out_of_range(self):
# sucet tipov > pocet kopiek -> zakazana hodnota by bola zaporna,
# ziadne dodatocne maskovanie
r = Round(0, 0)
r.add_player_guess(0, 8)
r.add_player_guess(1, 5)
r.add_player_guess(2, 0)
mask = guess_mask(r)
self.assertEqual(mask, [True] * 9)
class MaskEngineConsistencyCase(unittest.TestCase):
"""Fuzz: maska presne zrkadli engine -- povolena akcia NIKDY nezlyha,
zakazana akcia VZDY vyhodi BridzikException."""
def _check_guess_mask(self, rnd, player, mask):
for g in range(N_GUESS_ACTIONS):
if mask[g]:
copy.deepcopy(rnd).add_player_guess(player, g)
else:
with self.assertRaises(BridzikException):
rnd.add_player_guess(player, g)
def _check_play_mask(self, rnd, player, mask):
self.assertIn(True, mask) # aktivny hrac ma vzdy legalny tah
for i in range(N_PLAY_ACTIONS):
if mask[i]:
copy.deepcopy(rnd).play_card(player, index_card(i))
else:
with self.assertRaises(BridzikException):
rnd.play_card(player, index_card(i))
def _check_observation(self, rnd, player):
obs = encode_observation(rnd, player)
self.assertEqual(len(obs), OBS_DIM)
for v in obs:
self.assertGreaterEqual(v, 0.0)
self.assertLessEqual(v, 1.0)
hand = {card_index(c) for c in rnd.player_cards[player]}
for i in range(N_CARDS):
self.assertEqual(obs[OFF_HAND + i], 1.0 if i in hand else 0.0)
if i in hand: # ruka a videne karty su disjunktne
self.assertEqual(obs[OFF_SEEN + i], 0.0)
# zakodovany void nikdy neprotireci realnej ruke hraca
for seat in range(4):
rel = relative_seat(seat, player)
held = {c.color for c in rnd.player_cards[seat]}
for ci, color in enumerate(COLORS):
if obs[OFF_VOIDS + rel * 4 + ci] == 1.0:
self.assertNotIn(color, held)
def _fuzz_round(self, rng, round_number, first_player):
rnd = Round(round_number, first_player)
for _ in range(4):
player = rnd.get_active_player()
self._check_observation(rnd, player)
mask = guess_mask(rnd)
self._check_guess_mask(rnd, player, mask)
rnd.add_player_guess(
player, rng.choice([g for g in range(N_GUESS_ACTIONS) if mask[g]])
)
while not rnd.is_completed():
player = rnd.get_active_player()
self._check_observation(rnd, player)
mask = play_mask(rnd, player)
self._check_play_mask(rnd, player, mask)
rnd.play_card(
player, index_card(rng.choice([i for i in range(N_PLAY_ACTIONS) if mask[i]]))
)
# kolo dohrane do konca cisto cez masky -> bodovanie funguje
self.assertEqual(len(rnd.get_points_summary()), 4)
def test_fuzz_all_round_numbers_and_seats(self):
rng = random.Random(1337)
for round_number in range(8):
for first_player in range(4):
for _ in range(3):
self._fuzz_round(rng, round_number, first_player)
if __name__ == '__main__':
unittest.main(verbosity=2)