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>
This commit is contained in:
@@ -0,0 +1,296 @@
|
||||
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)
|
||||
Reference in New Issue
Block a user