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>
297 lines
13 KiB
Python
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)
|