Files
bridzik/rl/encoding.py
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

169 lines
6.6 KiB
Python

"""Observation a action-mask encoding nad `Round` objektom z `bridzik.py`.
Ciste funkcie bez externych zavislosti (ziadny numpy/torch) -- vystupy su
obycajne zoznamy floatov/boolov, konverzia na tenzory je vecou volajuceho.
Tvrdy invariant (viz rl/DESIGN.md): VSETKY 4-hracske zlozky observacie su
egocentricky rotovane -- "ja" (parameter `player`) je vzdy index 0, ostatni
hraci +1/+2/+3 v smere hry. Bez tejto rotacie parameter sharing medzi
sedadlami nefunguje.
Masky zrkadlia pravidla enginu: maska NIKDY nesmie povolit akciu, ktoru by
`Round.add_player_guess` / `Round.play_card` odmietli, a naopak. Tuto zhodu
vynucuje fuzz-test v tests/test_encoding.py.
"""
from bridzik import Card, Card_colors, Card_values, ROUNDS_PER_SERIES
# Kanonicke poradie farieb a hodnot = poradie deklaracie v enum-och.
COLORS = list(Card_colors) # HEARTS, LEAVES, ACORNS, BELLS
VALUES = list(Card_values) # C7, C8, C9, C10, LOWER, UPPER, KING, ACE
N_CARDS = 32
N_GUESS_ACTIONS = 9 # tipy 0..8
N_PLAY_ACTIONS = N_CARDS # jedna akcia = jedna karta
_COLOR_INDEX = {color: i for i, color in enumerate(COLORS)}
_VALUE_INDEX = {value: i for i, value in enumerate(VALUES)}
# Layout observacie -- offsety su sucast verejneho kontraktu (testy aj siet
# sa na ne odkazuju menom, nie magickym cislom).
OFF_HAND = 0 # 32 multi-hot: vlastna ruka
OFF_SEEN = OFF_HAND + N_CARDS # 32 multi-hot: karty z dokoncenych kopiek
OFF_ROUND = OFF_SEEN + N_CARDS # 8 one-hot: round_number
OFF_PHASE = OFF_ROUND + ROUNDS_PER_SERIES # 1 flag: 1.0 = tipovacia faza
OFF_GUESSES = OFF_PHASE + 1 # 4 x (flag "uz tipoval", tip/8), rel. poradie
OFF_TRICKS = OFF_GUESSES + 8 # 4 x (vyhrane kopky / 8), rel. poradie
OFF_STASH = OFF_TRICKS + 4 # 4 x 32 one-hot: aktualna kopka, rel. poradie
OFF_STASH_LEADER = OFF_STASH + 4 * N_CARDS # 4 one-hot: rel. first_player kopky
OFF_VOIDS = OFF_STASH_LEADER + 4 # 4 hraci x 4 farby: dedukovane voidy, rel. poradie
OBS_DIM = OFF_VOIDS + 16 # = 233
def card_index(card: Card) -> int:
"""Index karty 0..31: farba (blok po 8) + hodnota."""
return _COLOR_INDEX[card.color] * 8 + _VALUE_INDEX[card.value]
def index_card(index: int) -> Card:
"""Inverzia card_index."""
return Card(COLORS[index // 8], VALUES[index % 8])
def relative_seat(seat: int, player: int) -> int:
"""Egocentricka rotacia: `player` -> 0, dalsi v smere hry -> 1, 2, 3."""
return (seat - player) % 4
def deduce_voids(rnd) -> dict:
"""Isto-dedukovane chybajuce farby hracov z priebehu kola.
Follow-suit pravidlo prezradza: kto nepriznal vynasanu farbu, uz ju nema;
kto pri tom nezahral ani cerven (tromf), nema ani tu. Vynasajuci hrac
neprezradza nic. Vracia dict seat -> set(Card_colors). Kedze karty pocas
kola len ubudaju, raz dedukovany void plati do konca kola.
"""
voids = {seat: set() for seat in range(4)}
for stash in rnd.stashes:
first = stash.get_first_card()
if first is None:
continue
for seat, card in stash.get_cards().items():
if seat == stash.first_player:
continue
if card.color != first.color:
voids[seat].add(first.color)
if card.color != Card_colors['HEARTS']:
voids[seat].add(Card_colors['HEARTS'])
return voids
def encode_observation(rnd, player: int) -> list:
"""Observacia kola z pohladu hraca `player` (OBS_DIM floatov v [0, 1]).
Funguje v oboch fazach (tipovanie aj hra) aj na terminalnom stave;
necita get_active_player(), takze sa da zavolat pre lubovolneho hraca
kedykolvek.
"""
obs = [0.0] * OBS_DIM
for card in rnd.player_cards[player]:
obs[OFF_HAND + card_index(card)] = 1.0
# Dokoncene kopky -> "videne karty"; jedina pripadna nedokoncena kopka
# (posledna) je aktualna rozohrana a koduje sa do slotov nizsie.
current_stash = None
for stash in rnd.stashes:
if stash.is_completed():
for card in stash.get_cards().values():
obs[OFF_SEEN + card_index(card)] = 1.0
else:
current_stash = stash
obs[OFF_ROUND + rnd.round_number] = 1.0
obs[OFF_PHASE] = 0.0 if rnd.is_guessing_completed() else 1.0
for seat, guess in rnd.guesses.items():
rel = relative_seat(seat, player)
obs[OFF_GUESSES + 2 * rel] = 1.0
obs[OFF_GUESSES + 2 * rel + 1] = guess / 8
tricks = rnd.get_stashes_winner_summary()
for seat in range(4):
obs[OFF_TRICKS + relative_seat(seat, player)] = tricks[seat] / 8
if current_stash is not None:
for seat, card in current_stash.get_cards().items():
rel = relative_seat(seat, player)
obs[OFF_STASH + rel * N_CARDS + card_index(card)] = 1.0
obs[OFF_STASH_LEADER + relative_seat(current_stash.first_player, player)] = 1.0
for seat, banned in deduce_voids(rnd).items():
rel = relative_seat(seat, player)
for color in banned:
obs[OFF_VOIDS + rel * 4 + _COLOR_INDEX[color]] = 1.0
return obs
def guess_mask(rnd) -> list:
"""Maska legalnych tipov (N_GUESS_ACTIONS boolov) pre aktivneho tipujuceho.
Zrkadli Round.add_player_guess: tip 0..(8 - round_number); poslednemu
(stvrtemu) tipujucemu je navyse zakazana hodnota, pri ktorej by sucet
tipov vysiel presne na pocet kopiek v kole.
"""
max_guess = 8 - rnd.round_number
mask = [g <= max_guess for g in range(N_GUESS_ACTIONS)]
if len(rnd.guesses) == 3:
forbidden = max_guess - sum(rnd.guesses.values())
if 0 <= forbidden <= max_guess:
mask[forbidden] = False
return mask
def legal_cards(hand: list, first_card) -> list:
"""Karty z `hand`, ktore smie hrac zahrat do kopky vynasanej `first_card`.
Zrkadli follow-suit logiku Round.play_card: povinna farba prvej karty
kopky; ak ju hrac nema, povinna cervena (tromf); ak nema ani tu,
lubovolna karta. Pri vynasani (`first_card is None`) lubovolna karta.
"""
if first_card is None:
return list(hand)
same_color = [c for c in hand if c.color == first_card.color]
if same_color:
return same_color
hearts = [c for c in hand if c.color == Card_colors['HEARTS']]
return hearts if hearts else list(hand)
def play_mask(rnd, player: int) -> list:
"""Maska legalnych kariet (N_PLAY_ACTIONS boolov) pre hraca `player`."""
stash = rnd.get_last_stash()
first_card = stash.get_first_card() if stash is not None else None
mask = [False] * N_PLAY_ACTIONS
for card in legal_cards(rnd.player_cards[player], first_card):
mask[card_index(card)] = True
return mask