Compare commits
17
Commits
9ade92d7fe
...
8e18325660
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8e18325660 | ||
|
|
973c279cbd | ||
|
|
b1010ae008 | ||
|
|
23de3ae3ac | ||
|
|
9a750756c5 | ||
|
|
8f2449a408 | ||
|
|
e1733f4943 | ||
|
|
3710a68e37 | ||
|
|
1b14ace2cd | ||
|
|
f17f85ebd9 | ||
|
|
558c8ec2d9 | ||
|
|
b1ca4bc78d | ||
|
|
6c95c64398 | ||
|
|
ca0cf413ee | ||
|
|
fb90737944 | ||
|
|
1fbba5a7e1 | ||
|
|
80929e7fed |
@@ -10,3 +10,6 @@ frontend
|
||||
.env
|
||||
.env.*
|
||||
docker-compose*.yaml
|
||||
# trenovacie artefakty do image nepatria (produkcia cita len rl/weights/)
|
||||
rl/checkpoints
|
||||
rl/runs
|
||||
|
||||
@@ -10,3 +10,5 @@ frontend/.vite/
|
||||
.env.*
|
||||
!.env.example
|
||||
geoip/*.mmdb
|
||||
rl/runs/
|
||||
rl/checkpoints/
|
||||
|
||||
@@ -11,6 +11,7 @@ RUN pip install --no-cache-dir -r requirements.txt
|
||||
COPY bridzik.py app.py ./
|
||||
COPY api ./api
|
||||
COPY db ./db
|
||||
COPY rl ./rl
|
||||
COPY tests ./tests
|
||||
|
||||
RUN useradd --create-home --uid 1000 appuser \
|
||||
|
||||
@@ -11,6 +11,10 @@ Hodinová routine v Claude Code číta túto sekciu, implementuje a presúva hot
|
||||
|
||||
## Hotovo
|
||||
|
||||
- [x] Skús upraviť dizajn scrollbaru v zozname bodov v hre tak, aby zodpovedal celej hre. — 2026-07-07: pridaný `.velvet-scroll` v `index.css` (tenký priehľadný track + zlatý polopriehľadný thumb, webkit aj Firefox `scrollbar-color`) a aplikovaný na scrollovateľný zoznam kôl v `Standings.tsx` (herný sidebar aj mobilný panel). Commit 6c95c64.
|
||||
- [x] V hre keď sa hádžu karty, pri hodení poslednej karty všetky karty na chvíľu zmiznú a potom sa znovu objavia cez animáciu. Uprav to tak, aby sa posledná karta pridala k predchádzajúcim bez zmiznutia kôpky, a až následne kôpka zmizla smerom k hráčovi, ktorý ju zobral. — 2026-07-07: `GameTable.tsx` číta `previous_stash` synchrónne (namiesto `lingeredStash` nastavovaného v `useEffect`), takže kôpka pri 4. karte nezmizne — posledná karta len pribudne k trom. Po krátkej pauze (`SETTLE_MS`) sa celá kôpka animáciou (`collect-*` keyframes v `index.css`) odsunie k sedadlu víťaza; víťaz sa počíta cez nový `stashWinner` v `gameRules.ts` podľa pravidiel enginu. Commit fb90737.
|
||||
- [x] Urob samostatné scrollovanie v zozname bodov v hlavnej hre, lebo teraz tam nevidno všetky hry. Zároveň dorob, aby sa po jednotlivých sériách dali body collapsnúť/expandnúť a ostali by len celkové body za sériu. — 2026-07-07: `Standings.tsx` rozdelený na fixnú hlavičku, scrollovateľný zoznam kôl (`overflow-y-auto`, na mobile `max-h-[45vh]`) a fixné súčty; dokončená séria (Σ riadok) je teraz klikateľná a zbaľuje/rozbaľuje svoje kolá. Commit 80929e7.
|
||||
- [x] V hre aj v histórii zruš preciarkavanie pri nesprávnych tipoch. — 2026-07-07: odstránené `line-through` z `Standings.tsx` (herný pohľad) aj `History.tsx` (detail hry) pri neúspešnom tipe. Commit 80929e7.
|
||||
- [x] V lobby hre sa dá skopírovať ID hry. Zmeň to tak, aby sa namiesto ID kopírovala celá URL linka na hru. — 2026-07-03: tlačidlo "Kopírovať" v Lobby.tsx teraz kopíruje `${window.location.origin}/lobby/${gid}` namiesto holého gid. Commit 49ac0d3.
|
||||
|
||||
<!-- Sem routine presunie dokončené úlohy s dátumom a krátkym popisom + commit hashom. -->
|
||||
|
||||
+179
-4
@@ -13,8 +13,10 @@ import socketio
|
||||
from bridzik import Bridzik, BridzikException, Card
|
||||
from db.db import init_db
|
||||
from api import auth as auth_module, history
|
||||
from api import bots as bots_module
|
||||
from api import stats as stats_module
|
||||
from api.auth import AuthError, RegistrationIncomplete
|
||||
from rl.encoding import index_card
|
||||
|
||||
|
||||
def _env_bool(name: str, default: bool) -> bool:
|
||||
@@ -179,6 +181,9 @@ class Game:
|
||||
self.players: list["Player"] = []
|
||||
self.started = False
|
||||
self.bridzik_core: Bridzik | None = None
|
||||
# Serializuje tahovu slucku botov -- dva sucasne _run_bot_turns tasky
|
||||
# by inak mohli tahat za to iste sedadlo.
|
||||
self.bot_lock = asyncio.Lock()
|
||||
|
||||
def start(self):
|
||||
self.bridzik_core = Bridzik()
|
||||
@@ -202,6 +207,10 @@ class Player:
|
||||
self.player_id = player_id # persistent account id (db.models.Player.id)
|
||||
self.token = str(uuid.uuid4()) # secret token used for secure reconnect
|
||||
self.connected = True
|
||||
# Bot = sedadlo bez socketu; `brain` je rozhodovaci objekt s rozhranim
|
||||
# guess(rnd, seat) / play(rnd, seat) z rl/players.py.
|
||||
self.is_bot = False
|
||||
self.brain = None
|
||||
|
||||
|
||||
class CardStatusEncoder(JSONEncoder):
|
||||
@@ -214,7 +223,13 @@ class CardStatusEncoder(JSONEncoder):
|
||||
|
||||
|
||||
def public_games() -> list:
|
||||
"""Public lobby view — no sids, no reconnect tokens."""
|
||||
"""Public lobby view — no sids, no reconnect tokens.
|
||||
|
||||
A game that finished naturally (all 4 series played out) stays in the
|
||||
`games` dict for reconnect purposes (e.g. a reload while still on the
|
||||
GameOver screen), but it has nothing left to offer the lobby — drop it
|
||||
here rather than have it linger forever as a "started"/resumable entry.
|
||||
"""
|
||||
return [
|
||||
{
|
||||
"gid": g.gid,
|
||||
@@ -226,11 +241,13 @@ def public_games() -> list:
|
||||
"name": p.name,
|
||||
"connected": p.connected,
|
||||
"player_id": p.player_id,
|
||||
"is_bot": p.is_bot,
|
||||
}
|
||||
for p in g.players
|
||||
],
|
||||
}
|
||||
for g in games.values()
|
||||
if g.bridzik_core is None or not g.bridzik_core.is_completed()
|
||||
]
|
||||
|
||||
|
||||
@@ -255,7 +272,8 @@ async def send_game_status(gid: str):
|
||||
"completed": core.is_completed(),
|
||||
# Self-contained roster so the game view doesn't depend on the lobby snapshot.
|
||||
"players": [
|
||||
{"order": p.order, "name": p.name, "connected": p.connected}
|
||||
{"order": p.order, "name": p.name, "connected": p.connected,
|
||||
"is_bot": p.is_bot}
|
||||
for p in sorted(game.players, key=lambda p: p.order)
|
||||
],
|
||||
"series_number": core.series[-1].series_number,
|
||||
@@ -292,11 +310,16 @@ async def _cleanup_abandoned_lobby(gid: str):
|
||||
task neskodny no-op (netreba nic rusit)."""
|
||||
await asyncio.sleep(LOBBY_ABANDON_GRACE_SECONDS)
|
||||
game = games.get(gid)
|
||||
if game is not None and not game.started and not any(p.connected for p in game.players):
|
||||
if game is not None and not game.started and not _any_human_connected(game):
|
||||
del games[gid]
|
||||
await broadcast_lobby()
|
||||
|
||||
|
||||
def _any_human_connected(game: "Game") -> bool:
|
||||
"""Boti su 'pripojeni' stale, pre opustenost lobby sa pocitaju len ludia."""
|
||||
return any(p.connected for p in game.players if not p.is_bot)
|
||||
|
||||
|
||||
async def _mark_player_offline(game: "Game", player: "Player"):
|
||||
"""Mark player disconnected. An unstarted game with nobody left gets a
|
||||
delayed cleanup (mobile sockets drop on screen lock, so an immediate
|
||||
@@ -304,7 +327,7 @@ async def _mark_player_offline(game: "Game", player: "Player"):
|
||||
is kept in memory so it stays in the lobby and can be resumed (it's torn
|
||||
down only by end_game)."""
|
||||
player.connected = False
|
||||
if not any(p.connected for p in game.players) and not game.started:
|
||||
if not _any_human_connected(game) and not game.started:
|
||||
asyncio.create_task(_cleanup_abandoned_lobby(game.gid))
|
||||
await sio.emit(
|
||||
"player_connection",
|
||||
@@ -345,11 +368,93 @@ def _load_game_into_memory(info: dict) -> "Game":
|
||||
for seat, (pid, uname) in enumerate(info["seats"]):
|
||||
player = Player(None, uname, seat, pid)
|
||||
player.connected = False
|
||||
# Boti sa rozpoznaju konvenciou mena a dostanu novy mozog -- ozivi ich
|
||||
# prvy _kick_bots (napr. ked sa clovek vrati cez rejoin_game).
|
||||
if bots_module.is_bot_username(uname):
|
||||
player.is_bot = True
|
||||
player.brain = bots_module.make_brain(uname)
|
||||
player.connected = True
|
||||
game.players.append(player)
|
||||
games[info["gid"]] = game
|
||||
return game
|
||||
|
||||
|
||||
# --- bot turns --------------------------------------------------------------
|
||||
|
||||
# Pauza medzi tahmi bota, nech ludia stihaju sledovat hru (0 = okamzite).
|
||||
BOT_MOVE_DELAY_SECONDS = float(os.environ.get("BOT_MOVE_DELAY_SECONDS", "0.8"))
|
||||
# Kopka na stole sa po dohrati este chvilu zmieta smerom k vitazovi (SETTLE_MS +
|
||||
# COLLECT_MS vo frontend/src/pages/GameTable.tsx, spolu 1650ms) -- kym tato
|
||||
# animacia nedobehne vsetkym hracom, prvy bot na tahu nesmie zahodit kartu do
|
||||
# novej kopky, inak by mu karta "vyletela" uprostred zmetania predoslej.
|
||||
TRICK_SWEEP_SECONDS = 1.7
|
||||
|
||||
|
||||
def _kick_bots(gid: str) -> None:
|
||||
"""Ak je v hre bot, spusti (na pozadi) dohratie vsetkych botich tahov.
|
||||
Vola sa po kazdej akcii, ktora mohla posunut tah na botie sedadlo."""
|
||||
game = games.get(gid)
|
||||
if game is not None and game.started and any(p.is_bot for p in game.players):
|
||||
asyncio.create_task(_run_bot_turns(gid))
|
||||
|
||||
|
||||
async def _run_bot_turns(gid: str):
|
||||
"""Kym je na tahu botie sedadlo, vykonavaj jeho tahy tym istym internym
|
||||
postupom ako handlery add_guess/play_card (engine validuje, historia sa
|
||||
zapisuje, room dostava game_status). MC vypocet bezi v executori, aby
|
||||
nedrzal event loop ostatnych hier."""
|
||||
game = games.get(gid)
|
||||
if game is None or not game.started or game.bridzik_core is None:
|
||||
return
|
||||
async with game.bot_lock:
|
||||
core = game.bridzik_core
|
||||
loop = asyncio.get_running_loop()
|
||||
while not core.is_completed():
|
||||
rnd = core.series[-1].get_last_round()
|
||||
seat = rnd.get_active_player()
|
||||
bot = game.player_by_order(seat)
|
||||
if bot is None or not bot.is_bot:
|
||||
return # na tahu je clovek
|
||||
delay = BOT_MOVE_DELAY_SECONDS
|
||||
if rnd.is_guessing_completed() and not rnd.get_last_stash().get_cards() \
|
||||
and core.get_previous_stash() is not None:
|
||||
# Bot vedie novu kopku a este bezi zmetanie tej predoslej.
|
||||
delay = max(delay, TRICK_SWEEP_SECONDS)
|
||||
if delay > 0:
|
||||
await asyncio.sleep(delay)
|
||||
if games.get(gid) is not game:
|
||||
return # hru medzitym niekto ukoncil (end_game)
|
||||
played_card = False
|
||||
try:
|
||||
if not rnd.is_guessing_completed():
|
||||
guess = await loop.run_in_executor(None, bot.brain.guess, rnd, seat)
|
||||
core.add_player_guess(seat, guess)
|
||||
else:
|
||||
action = await loop.run_in_executor(None, bot.brain.play, rnd, seat)
|
||||
core.play_card(seat, index_card(action))
|
||||
await history.record_completed_rounds(gid, core)
|
||||
played_card = True
|
||||
except BridzikException as exc:
|
||||
# Nemalo by nastat (bot hra len legalne tahy) -- nezacykli sa,
|
||||
# slucku znovu spusti dalsia akcia cloveka.
|
||||
await send_error_room(gid, str(exc))
|
||||
return
|
||||
# game_status musi ist PRED player_cards -- klient podla neho (novy
|
||||
# round_number + previous_stash) pozna, ze prave zacalo nove kolo, a
|
||||
# dovtedy si drzi starych karty na obrazovke (pozri "displayedHand" vo
|
||||
# frontend/src/pages/GameTable.tsx), kym nedobehne animacia zmetenia
|
||||
# poslednej kopky. Opacne poradie by novu ruku odhalilo predcasne.
|
||||
await send_game_status(gid)
|
||||
if played_card:
|
||||
for player in game.players:
|
||||
if player.sid:
|
||||
await send_player_cards(gid, player.order, player.sid)
|
||||
|
||||
|
||||
async def send_error_room(gid: str, message: str):
|
||||
await sio.emit("error", {"error": message}, room=gid)
|
||||
|
||||
|
||||
# --- connection lifecycle -------------------------------------------------
|
||||
|
||||
@sio.event
|
||||
@@ -481,6 +586,62 @@ async def register_player(sid, gid):
|
||||
await broadcast_lobby()
|
||||
|
||||
|
||||
@sio.on("add_bot")
|
||||
async def add_bot(sid, gid, kind=None):
|
||||
"""Hostitel prida bota na najnizsie volne sedadlo nezacatej hry."""
|
||||
sess = sessions.get(sid)
|
||||
if sess is None or sess["gid"] != gid:
|
||||
return await send_error(sid, "Nie ste v tejto hre.")
|
||||
if sess["order"] != 0:
|
||||
return await send_error(sid, "Iba hostitel moze pridavat botov.")
|
||||
game = games.get(gid)
|
||||
if game is None:
|
||||
return await send_error(sid, "Hra neexistuje.")
|
||||
if game.started:
|
||||
return await send_error(sid, "Hra uz zacala.")
|
||||
if len(game.players) >= 4:
|
||||
return await send_error(sid, "Prekroceny pocet hracov.")
|
||||
|
||||
if kind == "neural" and not bots_module.neural_available():
|
||||
return await send_error(sid, "AI bot nie je na tomto serveri dostupny.")
|
||||
if kind not in bots_module.BOT_KINDS:
|
||||
kind = bots_module.DEFAULT_KIND
|
||||
account = await bots_module.ensure_bot_account(
|
||||
kind, {p.player_id for p in game.players}
|
||||
)
|
||||
used = {p.order for p in game.players}
|
||||
order = next(o for o in range(4) if o not in used)
|
||||
player = Player(None, account["username"], order, account["player_id"])
|
||||
player.is_bot = True
|
||||
player.brain = bots_module.make_brain(account["username"])
|
||||
game.players.append(player)
|
||||
await broadcast_lobby()
|
||||
|
||||
|
||||
@sio.on("remove_bot")
|
||||
async def remove_bot(sid, gid, order):
|
||||
"""Hostitel odoberie bota z nezacatej hry (sedadlo sa uvolni)."""
|
||||
sess = sessions.get(sid)
|
||||
if sess is None or sess["gid"] != gid:
|
||||
return await send_error(sid, "Nie ste v tejto hre.")
|
||||
if sess["order"] != 0:
|
||||
return await send_error(sid, "Iba hostitel moze odoberat botov.")
|
||||
game = games.get(gid)
|
||||
if game is None:
|
||||
return await send_error(sid, "Hra neexistuje.")
|
||||
if game.started:
|
||||
return await send_error(sid, "Hra uz zacala.")
|
||||
try:
|
||||
seat = int(order)
|
||||
except (TypeError, ValueError):
|
||||
return await send_error(sid, "Neplatne sedadlo.")
|
||||
player = game.player_by_order(seat)
|
||||
if player is None or not player.is_bot:
|
||||
return await send_error(sid, "Na tomto sedadle nie je bot.")
|
||||
game.players.remove(player)
|
||||
await broadcast_lobby()
|
||||
|
||||
|
||||
@sio.on("leave_game")
|
||||
async def leave_game(sid):
|
||||
"""Explicit exit (e.g. a 'Back to lobby' button). The socket stays
|
||||
@@ -528,7 +689,11 @@ async def start_game(sid, gid):
|
||||
await broadcast_lobby()
|
||||
await send_game_status(gid)
|
||||
for player in game.players:
|
||||
# sid None = bot alebo offline sedadlo -- emit s to=None by karty
|
||||
# broadcastol VSETKYM klientom, preto sa preskakuje.
|
||||
if player.sid:
|
||||
await send_player_cards(gid, player.order, player.sid)
|
||||
_kick_bots(gid)
|
||||
|
||||
|
||||
@sio.on("end_game")
|
||||
@@ -583,6 +748,7 @@ async def reconnect_to_game(sid, gid, token):
|
||||
await send_player_cards(gid, player.order, sid)
|
||||
await sio.emit("player_connection", {"order": player.order, "connected": True}, room=gid)
|
||||
await broadcast_lobby()
|
||||
_kick_bots(gid)
|
||||
|
||||
|
||||
@sio.on("rejoin_game")
|
||||
@@ -620,6 +786,7 @@ async def rejoin_game(sid, gid):
|
||||
await send_player_cards(gid, player.order, sid)
|
||||
await sio.emit("player_connection", {"order": player.order, "connected": True}, room=gid)
|
||||
await broadcast_lobby()
|
||||
_kick_bots(gid)
|
||||
|
||||
|
||||
@sio.on("restore_game")
|
||||
@@ -678,6 +845,7 @@ async def add_guess(sid, guess):
|
||||
except BridzikException as exc:
|
||||
return await send_error(sid, str(exc))
|
||||
await send_game_status(game.gid)
|
||||
_kick_bots(game.gid)
|
||||
|
||||
|
||||
@sio.on("play_card")
|
||||
@@ -703,7 +871,14 @@ async def play_card(sid, card_key):
|
||||
await history.record_completed_rounds(game.gid, core)
|
||||
await send_game_status(game.gid)
|
||||
for player in game.players:
|
||||
if player.sid: # None (bot/offline) by broadcastoval karty vsetkym
|
||||
await send_player_cards(game.gid, player.order, player.sid)
|
||||
# A naturally-finished game (all 4 series played out) has nothing left to
|
||||
# offer the lobby -- refresh the list so it drops out immediately instead
|
||||
# of lingering as "started"/resumable until someone happens to leave it.
|
||||
if core.is_completed():
|
||||
await broadcast_lobby()
|
||||
_kick_bots(game.gid)
|
||||
|
||||
|
||||
# --- history (read-only) --------------------------------------------------
|
||||
|
||||
+120
@@ -0,0 +1,120 @@
|
||||
"""In-process boti: DB ucty botov a ich rozhodovacie "mozgy" z rl/players.py.
|
||||
|
||||
Bot je normalny hrac na sedadle -- ma riadok v tabulke `players` (aby
|
||||
historia, standings a restore fungovali bez zmeny), ale ziadny socket.
|
||||
Tahovu slucku botov ma api/__init__.py (_run_bot_turns); tu je len to,
|
||||
co potrebuje DB a rl vrstvu.
|
||||
|
||||
Bezpecnost botich uctov: totp_secret je nahodny a nikde sa neuklada v
|
||||
citatelnej podobe, totp_last_step sa nastavi na aktualny krok -- ucet tym
|
||||
padom NIE JE "nedokoncena registracia" (viz auth._is_unconfirmed), takze
|
||||
register_account ho odmietne prepisat a login bez secretu neprejde.
|
||||
"""
|
||||
|
||||
import os
|
||||
from random import Random
|
||||
|
||||
import pyotp
|
||||
from sqlalchemy import select
|
||||
|
||||
from api import auth
|
||||
from db import crypto
|
||||
from db.db import async_session
|
||||
from db.models import Player
|
||||
from rl.players import HeuristicPlayer, RandomPlayer
|
||||
from rl.pure_net import DEFAULT_WEIGHTS_PATH, PureNet, PureNeuralPlayer
|
||||
|
||||
# Prefix je konvencia na rozpoznanie bota (aj po restarte servera, kedy sa
|
||||
# sedadla obnovuju z DB len ako (player_id, username)).
|
||||
BOT_PREFIX = "bot:"
|
||||
DEFAULT_KIND = "heuristic"
|
||||
|
||||
# Natrenovana siet -- vahy (rl/weights/, export z rl/export.py) sa nacitaju
|
||||
# raz a zdielaju medzi botmi (PureNet je bezstavovy, len cita).
|
||||
_pure_net: PureNet | None = None
|
||||
|
||||
|
||||
def _neural_brain() -> PureNeuralPlayer:
|
||||
global _pure_net
|
||||
if _pure_net is None:
|
||||
_pure_net = PureNet.load()
|
||||
return PureNeuralPlayer(_pure_net)
|
||||
|
||||
|
||||
def neural_available() -> bool:
|
||||
return os.path.exists(DEFAULT_WEIGHTS_PATH)
|
||||
|
||||
|
||||
_BRAINS = {
|
||||
"heuristic": lambda: HeuristicPlayer(Random()),
|
||||
"random": lambda: RandomPlayer(Random()),
|
||||
"neural": _neural_brain,
|
||||
}
|
||||
BOT_KINDS = tuple(_BRAINS)
|
||||
|
||||
|
||||
def available_kinds() -> tuple:
|
||||
"""Druhy botov ponuknutelne na tomto serveri (neural len s vahami)."""
|
||||
return tuple(k for k in _BRAINS if k != "neural" or neural_available())
|
||||
|
||||
|
||||
def is_bot_username(username: str) -> bool:
|
||||
return bool(username) and username.startswith(BOT_PREFIX)
|
||||
|
||||
|
||||
def kind_of(username: str) -> str:
|
||||
"""'bot:heuristic-2' -> 'heuristic'; neznamy druh padne na DEFAULT_KIND."""
|
||||
body = username[len(BOT_PREFIX):]
|
||||
kind = body.rsplit("-", 1)[0]
|
||||
return kind if kind in _BRAINS else DEFAULT_KIND
|
||||
|
||||
|
||||
def make_brain(username: str):
|
||||
"""Rozhodovaci objekt (guess/play rozhranie z rl/players.py) pre bota.
|
||||
|
||||
Neural bez suboru vah (napr. restore hry na serveri bez exportu) padne
|
||||
na heuristiku -- sedadlo hra dalej, len inym mozgom.
|
||||
"""
|
||||
kind = kind_of(username)
|
||||
if kind == "neural" and not neural_available():
|
||||
kind = DEFAULT_KIND
|
||||
return _BRAINS[kind]()
|
||||
|
||||
|
||||
def _suffix_number(username: str) -> int:
|
||||
try:
|
||||
return int(username.rsplit("-", 1)[1])
|
||||
except (IndexError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
async def ensure_bot_account(kind: str, exclude_ids: set) -> dict:
|
||||
"""Najde alebo zalozi boti ucet daneho druhu; vrati {player_id, username}.
|
||||
|
||||
`exclude_ids` su ucty uz obsadene v danej hre -- kazde sedadlo potrebuje
|
||||
INY ucet (Game.playerN_id aj unikat v Guess predpokladaju 4 rozne ID).
|
||||
Ucty sa cisluju bot:<kind>-1, -2, ... a recykluju sa medzi hrami.
|
||||
"""
|
||||
prefix = f"{BOT_PREFIX}{kind}-"
|
||||
async with async_session() as session:
|
||||
rows = (
|
||||
await session.scalars(
|
||||
select(Player).where(Player.username.like(prefix + "%"))
|
||||
)
|
||||
).all()
|
||||
for player in sorted(rows, key=lambda p: _suffix_number(p.username)):
|
||||
if player.id not in exclude_ids:
|
||||
return {"player_id": player.id, "username": player.username}
|
||||
|
||||
number = 1 + max((_suffix_number(p.username) for p in rows), default=0)
|
||||
username = f"{prefix}{number}"
|
||||
player = Player(
|
||||
username=username,
|
||||
# nahodny secret, ktory sa zahodi -- nikto sa zan neprihlasi
|
||||
totp_secret=crypto.encrypt(pyotp.random_base32()),
|
||||
# nenulovy last_step = ucet sa netvari ako nedokoncena registracia
|
||||
totp_last_step=auth._current_step(),
|
||||
)
|
||||
session.add(player)
|
||||
await session.commit()
|
||||
return {"player_id": player.id, "username": player.username}
|
||||
+12
@@ -19,6 +19,10 @@ class Card_colors(Enum):
|
||||
return self.name == other.name
|
||||
return NotImplemented
|
||||
|
||||
# vlastne __eq__ rusi zdedeny __hash__ -- obnovit konzistentne s __eq__
|
||||
def __hash__(self):
|
||||
return hash(self.name)
|
||||
|
||||
|
||||
class Card_values(Enum):
|
||||
C7 = 1
|
||||
@@ -51,6 +55,10 @@ class Card_values(Enum):
|
||||
return self.name == other.name
|
||||
return NotImplemented
|
||||
|
||||
# vlastne __eq__ rusi zdedeny __hash__ -- obnovit konzistentne s __eq__
|
||||
def __hash__(self):
|
||||
return hash(self.name)
|
||||
|
||||
|
||||
class Card():
|
||||
def __init__(self, color: Card_colors, value: Card_values):
|
||||
@@ -63,6 +71,10 @@ class Card():
|
||||
and self.value == other.value
|
||||
return NotImplemented
|
||||
|
||||
# vlastne __eq__ rusi zdedeny __hash__ -- obnovit konzistentne s __eq__
|
||||
def __hash__(self):
|
||||
return hash((self.color, self.value))
|
||||
|
||||
def __str__(self):
|
||||
return '{}_{}'.format(self.color.name, self.value.name)
|
||||
|
||||
|
||||
@@ -26,7 +26,6 @@ interface Props {
|
||||
|
||||
export default function Hand({ hand, myTurn, isPlayPhase, playableKeys, desktop = false }: Props) {
|
||||
const groups = groupedByColor(hand);
|
||||
if (groups.length === 0) return null;
|
||||
|
||||
const canPlay = isPlayPhase && myTurn;
|
||||
|
||||
@@ -52,6 +51,10 @@ export default function Hand({ hand, myTurn, isPlayPhase, playableKeys, desktop
|
||||
<div className="h-px flex-1 max-w-[80px] bg-gradient-to-l from-transparent to-gold/20" />
|
||||
</div>
|
||||
|
||||
{/* Fixed height reserves the card row even when the hand is briefly
|
||||
empty (last card of a round just played, next deal not in yet) --
|
||||
otherwise this whole area collapses and the layout jumps. */}
|
||||
<div className="flex items-end justify-center" style={{ minHeight: desktop ? 100 : 84 }}>
|
||||
{desktop ? (
|
||||
// Desktop has room — keep cards grouped by suit, wrap if needed.
|
||||
<div className="flex flex-wrap gap-3 justify-center items-end">
|
||||
@@ -67,6 +70,7 @@ export default function Hand({ hand, myTurn, isPlayPhase, playableKeys, desktop
|
||||
<MobileHand groups={groups} cardProps={cardProps} />
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
import { useEffect, useRef, useState } from 'react';
|
||||
|
||||
interface Props {
|
||||
username?: string;
|
||||
onHistory: () => void;
|
||||
onDonate: () => void;
|
||||
onLogout: () => void;
|
||||
}
|
||||
|
||||
/** Header navigation: a row of links from `sm:` up, a hamburger dropdown below it. */
|
||||
export default function HeaderMenu({ username, onHistory, onDonate, onLogout }: Props) {
|
||||
const [open, setOpen] = useState(false);
|
||||
const ref = useRef<HTMLDivElement>(null);
|
||||
|
||||
useEffect(() => {
|
||||
if (!open) return;
|
||||
const onPointerDown = (e: PointerEvent) => {
|
||||
if (!ref.current?.contains(e.target as Node)) setOpen(false);
|
||||
};
|
||||
const onKeyDown = (e: KeyboardEvent) => {
|
||||
if (e.key === 'Escape') setOpen(false);
|
||||
};
|
||||
document.addEventListener('pointerdown', onPointerDown);
|
||||
document.addEventListener('keydown', onKeyDown);
|
||||
return () => {
|
||||
document.removeEventListener('pointerdown', onPointerDown);
|
||||
document.removeEventListener('keydown', onKeyDown);
|
||||
};
|
||||
}, [open]);
|
||||
|
||||
const pick = (action: () => void) => () => {
|
||||
setOpen(false);
|
||||
action();
|
||||
};
|
||||
|
||||
return (
|
||||
<>
|
||||
<div className="hidden sm:flex items-center gap-3 text-sm">
|
||||
<span className="text-green-dim">{username}</span>
|
||||
<button onClick={onHistory} className="text-gold hover:text-gold-bright">
|
||||
História
|
||||
</button>
|
||||
<button onClick={onDonate} className="text-gold hover:text-gold-bright">
|
||||
Na kávu
|
||||
</button>
|
||||
<button onClick={onLogout} className="text-green-dim hover:text-gold">
|
||||
Odhlásiť
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div ref={ref} className="relative sm:hidden">
|
||||
<button
|
||||
onClick={() => setOpen((o) => !o)}
|
||||
aria-label="Menu"
|
||||
aria-expanded={open}
|
||||
aria-haspopup="menu"
|
||||
className="flex flex-col justify-center gap-[5px] w-9 h-9 items-center rounded-lg border border-gold/20 text-gold hover:border-gold/50 transition-colors"
|
||||
>
|
||||
<span className="block w-4 h-px bg-current" />
|
||||
<span className="block w-4 h-px bg-current" />
|
||||
<span className="block w-4 h-px bg-current" />
|
||||
</button>
|
||||
|
||||
{open && (
|
||||
<div
|
||||
role="menu"
|
||||
className="absolute right-0 top-11 z-40 w-44 bg-header border border-[#142018] rounded-xl py-1 shadow-[0_20px_60px_rgba(0,0,0,.6)]"
|
||||
>
|
||||
{username && (
|
||||
<p className="px-4 py-2 text-xs text-green-dim border-b border-[#142018] truncate">
|
||||
{username}
|
||||
</p>
|
||||
)}
|
||||
<button
|
||||
role="menuitem"
|
||||
onClick={pick(onHistory)}
|
||||
className="block w-full text-left px-4 py-2.5 text-sm text-gold hover:bg-gold/10"
|
||||
>
|
||||
História
|
||||
</button>
|
||||
<button
|
||||
role="menuitem"
|
||||
onClick={pick(onDonate)}
|
||||
className="block w-full text-left px-4 py-2.5 text-sm text-gold hover:bg-gold/10"
|
||||
>
|
||||
Na kávu
|
||||
</button>
|
||||
<button
|
||||
role="menuitem"
|
||||
onClick={pick(onLogout)}
|
||||
className="block w-full text-left px-4 py-2.5 text-sm text-green-dim hover:bg-gold/10 hover:text-gold"
|
||||
>
|
||||
Odhlásiť
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</>
|
||||
);
|
||||
}
|
||||
@@ -12,25 +12,34 @@ interface Props {
|
||||
export default function PlayerCircle({ name, won, guess, active, size = 52 }: Props) {
|
||||
const nameFont = Math.max(9, Math.round(size * 0.17));
|
||||
const valueFont = Math.round(size * 0.32);
|
||||
// Oval: width = size, height a touch shorter so it reads as an ellipse.
|
||||
// Oval: height a touch shorter than size so it reads as an ellipse. Width
|
||||
// starts at `size` (a circle for short names) but grows with the name via
|
||||
// fit-content + padding, up to a cap beyond which the name is ellipsised
|
||||
// rather than wrapping (wrapping would break the fixed height/oval shape).
|
||||
const height = Math.round(size * 0.78);
|
||||
const hPad = Math.round(size * 0.16);
|
||||
const maxWidth = Math.round(size * 2);
|
||||
|
||||
return (
|
||||
<div
|
||||
// Only the active player is highlighted (gold ring + glow) — colors come
|
||||
// from the velvet-table palette tokens (tailwind.config.js), not literals.
|
||||
className={`flex flex-col items-center justify-center rounded-[50%] ${
|
||||
className={`flex flex-col items-center justify-center rounded-full ${
|
||||
active ? 'bg-circle-active border-2 border-gold' : 'bg-circle border-[1.5px] border-gold/20'
|
||||
}`}
|
||||
style={{
|
||||
width: size,
|
||||
width: 'fit-content',
|
||||
minWidth: size,
|
||||
maxWidth,
|
||||
height,
|
||||
paddingLeft: hPad,
|
||||
paddingRight: hPad,
|
||||
boxShadow: '0 2px 10px rgba(0,0,0,.45)',
|
||||
animation: active ? 'ar 2.2s ease-in-out infinite' : undefined,
|
||||
}}
|
||||
>
|
||||
<span
|
||||
className={`uppercase leading-tight text-center ${active ? 'text-gold' : 'text-green-circle'}`}
|
||||
className={`uppercase leading-tight text-center overflow-hidden text-ellipsis whitespace-nowrap max-w-full ${active ? 'text-gold' : 'text-green-circle'}`}
|
||||
style={{
|
||||
fontFamily: '"DM Sans",sans-serif',
|
||||
fontSize: nameFont,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { useState } from 'react';
|
||||
import type { PlayerInfo } from '../types';
|
||||
import { computeTotal } from '../lib/standings';
|
||||
import { displayName } from '../lib/names';
|
||||
|
||||
interface Props {
|
||||
standings: number[][][];
|
||||
@@ -14,6 +15,15 @@ interface Props {
|
||||
|
||||
export default function Standings({ standings, guesses = [], players, myOrder, desktop = false }: Props) {
|
||||
const [open, setOpen] = useState(false);
|
||||
// Finished series (8/8 rounds) collapsed to just their Σ row; toggled per series.
|
||||
const [collapsedSeries, setCollapsedSeries] = useState<Set<number>>(new Set());
|
||||
const toggleSeries = (si: number) =>
|
||||
setCollapsedSeries((prev) => {
|
||||
const next = new Set(prev);
|
||||
if (next.has(si)) next.delete(si);
|
||||
else next.add(si);
|
||||
return next;
|
||||
});
|
||||
|
||||
// Player columns in seat order; the local player's column is highlighted.
|
||||
const cols = [...players].sort((a, b) => a.order - b.order);
|
||||
@@ -39,13 +49,11 @@ export default function Standings({ standings, guesses = [], players, myOrder, d
|
||||
total: desktop ? 20 : 18,
|
||||
};
|
||||
|
||||
const table = (
|
||||
<div className="flex-1 flex flex-col px-3 pt-3 pb-4">
|
||||
{/* Column headers */}
|
||||
<div
|
||||
className="grid items-end mb-1"
|
||||
style={{ gridTemplateColumns: `28px repeat(${cols.length}, 1fr)` }}
|
||||
>
|
||||
const gridCols = { gridTemplateColumns: `28px repeat(${cols.length}, 1fr)` };
|
||||
|
||||
const columnHeader = (
|
||||
<div className="px-3 pt-3 flex-shrink-0">
|
||||
<div className="grid items-end mb-1" style={gridCols}>
|
||||
<div />
|
||||
{cols.map((p) => (
|
||||
<div
|
||||
@@ -55,32 +63,35 @@ export default function Standings({ standings, guesses = [], players, myOrder, d
|
||||
}`}
|
||||
style={{ fontSize: fz.head }}
|
||||
>
|
||||
{p.name}
|
||||
{displayName(p.name)}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
<div className="h-px bg-gold/10 mb-1" />
|
||||
<div className="h-px bg-gold/10" />
|
||||
</div>
|
||||
);
|
||||
|
||||
{/* Completed rounds, grouped by series with a per-series summary row */}
|
||||
{standings.flatMap((seriesRounds, si) => {
|
||||
// Completed rounds, grouped by series with a per-series summary row. Finished
|
||||
// series can be collapsed to just their Σ row so long games stay scannable.
|
||||
const rows = standings.flatMap((seriesRounds, si) => {
|
||||
const priorRounds = seriesRoundOffsets[si];
|
||||
const elems = seriesRounds.map((scores, lri) => (
|
||||
<div
|
||||
key={`r-${si}-${lri}`}
|
||||
className="grid items-center py-1 border-b border-gold/[.05]"
|
||||
style={{ gridTemplateColumns: `28px repeat(${cols.length}, 1fr)` }}
|
||||
>
|
||||
const isFinished = seriesRounds.length === ROUNDS_PER_SERIES;
|
||||
const isCollapsed = isFinished && collapsedSeries.has(si);
|
||||
const elems = isCollapsed
|
||||
? []
|
||||
: seriesRounds.map((scores, lri) => (
|
||||
<div key={`r-${si}-${lri}`} className="grid items-center py-1 border-b border-gold/[.05]" style={gridCols}>
|
||||
<div className="text-center text-[#7a7252]" style={{ fontSize: fz.idx }}>
|
||||
{priorRounds + lri + 1}
|
||||
</div>
|
||||
{cols.map((p) => {
|
||||
const points = scores[p.order] ?? 0;
|
||||
// Failed tip (0 points) → show the struck-through tip instead of 0.
|
||||
// Failed tip → show the tip in the dimmer "0 points" color, no strikethrough.
|
||||
if (points === 0) {
|
||||
return (
|
||||
<div
|
||||
key={p.order}
|
||||
className="text-center font-serif leading-none line-through"
|
||||
className="text-center font-serif leading-none"
|
||||
style={{ fontSize: fz.cell, color: '#7a6e4a' }}
|
||||
>
|
||||
{guesses[si]?.[lri]?.[p.order] ?? 0}
|
||||
@@ -100,15 +111,21 @@ export default function Standings({ standings, guesses = [], players, myOrder, d
|
||||
</div>
|
||||
));
|
||||
|
||||
// After a finished series, sum its points per player.
|
||||
if (seriesRounds.length === ROUNDS_PER_SERIES) {
|
||||
// After a finished series, sum its points per player. Clicking toggles
|
||||
// whether that series' individual rounds are shown.
|
||||
if (isFinished) {
|
||||
elems.push(
|
||||
<div
|
||||
key={`s-${si}`}
|
||||
className="grid items-center py-1 my-0.5 rounded bg-gold/[.07]"
|
||||
style={{ gridTemplateColumns: `28px repeat(${cols.length}, 1fr)` }}
|
||||
role="button"
|
||||
tabIndex={0}
|
||||
onClick={() => toggleSeries(si)}
|
||||
onKeyDown={(e) => (e.key === 'Enter' || e.key === ' ') && toggleSeries(si)}
|
||||
className="grid items-center py-1 my-0.5 rounded bg-gold/[.07] cursor-pointer select-none"
|
||||
style={gridCols}
|
||||
>
|
||||
<div className="text-center font-serif text-gold" style={{ fontSize: fz.sigma }}>
|
||||
<div className="text-center font-serif text-gold flex items-center justify-center gap-[2px]" style={{ fontSize: fz.sigma }}>
|
||||
<span className="text-[8px] text-green-dim">{isCollapsed ? '▸' : '▾'}</span>
|
||||
Σ{si + 1}
|
||||
</div>
|
||||
{cols.map((p) => {
|
||||
@@ -129,27 +146,26 @@ export default function Standings({ standings, guesses = [], players, myOrder, d
|
||||
);
|
||||
}
|
||||
return elems;
|
||||
})}
|
||||
});
|
||||
|
||||
const scrollableRounds = (
|
||||
<div className={`velvet-scroll flex-1 min-h-0 overflow-y-auto px-3 ${desktop ? '' : 'max-h-[45vh]'}`}>
|
||||
{rows}
|
||||
|
||||
{/* Active round placeholder */}
|
||||
<div
|
||||
className="grid items-center py-1 rounded mt-0.5 bg-gold/[.04]"
|
||||
style={{ gridTemplateColumns: `28px repeat(${cols.length}, 1fr)` }}
|
||||
>
|
||||
<div className="grid items-center py-1 rounded mt-0.5 bg-gold/[.04]" style={gridCols}>
|
||||
<div className="text-center font-medium text-gold" style={{ fontSize: fz.idx }}>{completedRounds + 1}</div>
|
||||
{cols.map((p) => (
|
||||
<div key={p.order} className="text-center text-[#7a7252]" style={{ fontSize: fz.dot }}>·</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
|
||||
<div className="flex-1 min-h-2" />
|
||||
const totalsBlock = (
|
||||
<div className="px-3 pb-4 pt-2 flex-shrink-0">
|
||||
<div className="h-px bg-gold/20 mb-2" />
|
||||
|
||||
{/* Totals */}
|
||||
<div
|
||||
className="grid items-center py-0.5"
|
||||
style={{ gridTemplateColumns: `28px repeat(${cols.length}, 1fr)` }}
|
||||
>
|
||||
<div className="grid items-center py-0.5" style={gridCols}>
|
||||
<div className="text-center uppercase tracking-[.08em] text-green-dim" style={{ fontSize: fz.sigma }}>
|
||||
Σ
|
||||
</div>
|
||||
@@ -168,13 +184,21 @@ export default function Standings({ standings, guesses = [], players, myOrder, d
|
||||
</div>
|
||||
);
|
||||
|
||||
const content = (
|
||||
<div className="flex-1 min-h-0 flex flex-col">
|
||||
{columnHeader}
|
||||
{scrollableRounds}
|
||||
{totalsBlock}
|
||||
</div>
|
||||
);
|
||||
|
||||
if (desktop) {
|
||||
return (
|
||||
<aside className="w-[268px] flex-shrink-0 bg-header border-l border-[#142018] flex flex-col">
|
||||
<div className="h-[58px] flex items-center gap-2 px-5 border-b border-[#14221a]">
|
||||
<div className="h-[58px] flex items-center gap-2 px-5 border-b border-[#14221a] flex-shrink-0">
|
||||
<span className="font-serif uppercase tracking-[.12em] text-[13px] text-gold">Skóre</span>
|
||||
</div>
|
||||
{table}
|
||||
{content}
|
||||
</aside>
|
||||
);
|
||||
}
|
||||
@@ -189,7 +213,7 @@ export default function Standings({ standings, guesses = [], players, myOrder, d
|
||||
<span>Skóre</span>
|
||||
<span className="text-green-dim">{open ? '▲' : '▼'}</span>
|
||||
</button>
|
||||
{open && table}
|
||||
{open && content}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import type { PlayerInfo, StashData } from '../types';
|
||||
import CardView from './CardView';
|
||||
import { displayName } from '../lib/names';
|
||||
|
||||
interface Props {
|
||||
stash: StashData | null;
|
||||
@@ -18,7 +19,7 @@ export default function Trick({ stash, players, myOrder }: Props) {
|
||||
: [];
|
||||
|
||||
const nameFor = (order: number) =>
|
||||
players.find((p) => p.order === order)?.name ?? '';
|
||||
displayName(players.find((p) => p.order === order)?.name);
|
||||
|
||||
const overlap = -16;
|
||||
const slotH = 80;
|
||||
|
||||
@@ -30,6 +30,27 @@
|
||||
}
|
||||
}
|
||||
|
||||
/* Slim, on-theme scrollbar (velvet felt + gold accent) for opt-in scroll areas
|
||||
such as the in-game score list, instead of the chunky OS default. */
|
||||
.velvet-scroll {
|
||||
scrollbar-width: thin;
|
||||
scrollbar-color: rgba(201, 168, 76, 0.35) transparent;
|
||||
}
|
||||
.velvet-scroll::-webkit-scrollbar {
|
||||
width: 7px;
|
||||
height: 7px;
|
||||
}
|
||||
.velvet-scroll::-webkit-scrollbar-track {
|
||||
background: transparent;
|
||||
}
|
||||
.velvet-scroll::-webkit-scrollbar-thumb {
|
||||
background: rgba(201, 168, 76, 0.28);
|
||||
border-radius: 4px;
|
||||
}
|
||||
.velvet-scroll::-webkit-scrollbar-thumb:hover {
|
||||
background: rgba(201, 168, 76, 0.5);
|
||||
}
|
||||
|
||||
/* Velvet-table animations (design handoff). Declared as raw CSS so they work
|
||||
both via Tailwind's animate-* utilities and inline `animation:` strings. */
|
||||
@keyframes tp {
|
||||
@@ -70,3 +91,21 @@
|
||||
from { opacity: 0; transform: translateX(110px) scale(0.82); }
|
||||
to { opacity: 1; transform: none; }
|
||||
}
|
||||
|
||||
/* A completed trick is swept off the table towards the seat that won it. */
|
||||
@keyframes collect-bottom {
|
||||
from { opacity: 1; transform: none; }
|
||||
to { opacity: 0; transform: translateY(150px) scale(0.66); }
|
||||
}
|
||||
@keyframes collect-top {
|
||||
from { opacity: 1; transform: none; }
|
||||
to { opacity: 0; transform: translateY(-150px) scale(0.66); }
|
||||
}
|
||||
@keyframes collect-left {
|
||||
from { opacity: 1; transform: none; }
|
||||
to { opacity: 0; transform: translateX(-190px) scale(0.66); }
|
||||
}
|
||||
@keyframes collect-right {
|
||||
from { opacity: 1; transform: none; }
|
||||
to { opacity: 0; transform: translateX(190px) scale(0.66); }
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
import type { CardColor, Hand } from '../types';
|
||||
import type { CardColor, CardValue, Hand, StashData } from '../types';
|
||||
|
||||
export function computePlayable(hand: Hand, ledColor: CardColor | null): Set<string> {
|
||||
const keys = Object.keys(hand);
|
||||
@@ -13,6 +13,35 @@ export function computePlayable(hand: Hand, ledColor: CardColor | null): Set<str
|
||||
return new Set(keys);
|
||||
}
|
||||
|
||||
const VALUE_ORDER: CardValue[] = ['C7', 'C8', 'C9', 'C10', 'LOWER', 'UPPER', 'KING', 'ACE'];
|
||||
|
||||
/** Seat that wins a completed 4-card trick — mirrors Stash.get_winner in the
|
||||
* engine: HEARTS (červeň) is the permanent trump, otherwise the highest card
|
||||
* of the led colour wins. Safe to call on a partial stash (returns the current
|
||||
* leader among the cards played so far). */
|
||||
export function stashWinner(stash: StashData): number {
|
||||
const led = stash.cards[String(stash.first_player)];
|
||||
if (!led) return stash.first_player;
|
||||
let winner = stash.first_player;
|
||||
let best = led;
|
||||
for (let i = 0; i < 4; i++) {
|
||||
const c = stash.cards[String(i)];
|
||||
if (!c) continue;
|
||||
if (c.color === led.color || c.color === 'HEARTS') {
|
||||
if (c.color === best.color) {
|
||||
if (VALUE_ORDER.indexOf(c.value) >= VALUE_ORDER.indexOf(best.value)) {
|
||||
best = c;
|
||||
winner = i;
|
||||
}
|
||||
} else if (c.color === 'HEARTS') {
|
||||
best = c;
|
||||
winner = i;
|
||||
}
|
||||
}
|
||||
}
|
||||
return winner;
|
||||
}
|
||||
|
||||
/** The bid the last guesser may not make: the four bids must not sum to the
|
||||
* number of tricks in the round (mirrors Round.add_player_guess in the engine).
|
||||
* Returns null while earlier players are still guessing. */
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
/** Display helpers for player names. Bot accounts follow the server-side
|
||||
* convention "bot:<kind>-<n>" (see api/bots.py) — render them as a short
|
||||
* friendly label instead of the raw username. */
|
||||
|
||||
const BOT_PREFIX = 'bot:';
|
||||
|
||||
export function isBotName(name?: string | null): boolean {
|
||||
return !!name && name.startsWith(BOT_PREFIX);
|
||||
}
|
||||
|
||||
/** "bot:heuristic-2" -> "Bot 2", "bot:neural-1" -> "AI bot 1"; other kinds
|
||||
* keep a suffix ("Bot 1 (random)"); non-bot names pass through unchanged. */
|
||||
export function displayName(name?: string | null): string {
|
||||
if (!name || !isBotName(name)) return name ?? '';
|
||||
const body = name.slice(BOT_PREFIX.length);
|
||||
const dash = body.lastIndexOf('-');
|
||||
const kind = dash > 0 ? body.slice(0, dash) : body;
|
||||
const num = dash > 0 ? body.slice(dash + 1) : '';
|
||||
if (kind === 'neural') return num ? `AI bot ${num}` : 'AI bot';
|
||||
const label = num ? `Bot ${num}` : 'Bot';
|
||||
return kind === 'heuristic' ? label : `${label} (${kind})`;
|
||||
}
|
||||
@@ -50,6 +50,10 @@ export const emit = {
|
||||
// Reopen a prematurely-ended game from history back into the lobby.
|
||||
restoreGame: (gid: string) => socket.emit('restore_game', gid),
|
||||
leaveGame: () => socket.emit('leave_game'),
|
||||
// Bots: host-only, before the game starts (seat picked by the server).
|
||||
// kind: 'heuristic' (default) | 'neural' (trained net) | 'random'
|
||||
addBot: (gid: string, kind: string = 'heuristic') => socket.emit('add_bot', gid, kind),
|
||||
removeBot: (gid: string, order: number) => socket.emit('remove_bot', gid, order),
|
||||
endGame: (gid: string) => socket.emit('end_game', gid),
|
||||
startGame: (gid: string) => socket.emit('start_game', gid),
|
||||
reconnectToGame: (gid: string, token: string) => socket.emit('reconnect_to_game', gid, token),
|
||||
|
||||
@@ -3,6 +3,7 @@ import { useNavigate } from 'react-router-dom';
|
||||
import { useGameStore } from '../store/gameStore';
|
||||
import { emit, socket, setAuthToken } from '../lib/socket';
|
||||
import { trackEvent } from '../lib/track';
|
||||
import HeaderMenu from '../components/HeaderMenu';
|
||||
import NameModal from '../components/NameModal';
|
||||
import RulesModal from '../components/RulesModal';
|
||||
|
||||
@@ -33,18 +34,12 @@ export default function GameList() {
|
||||
<div className="max-w-md mx-auto p-4 pt-8 min-h-screen">
|
||||
<div className="flex items-center justify-between mb-6">
|
||||
<h1 className="font-serif text-2xl text-gold tracking-wide">Bridžik</h1>
|
||||
<div className="flex items-center gap-3 text-sm">
|
||||
<span className="text-green-dim">{account?.username}</span>
|
||||
<button onClick={() => navigate('/history')} className="text-gold hover:text-gold-bright">
|
||||
História
|
||||
</button>
|
||||
<button onClick={() => navigate('/donate')} className="text-gold hover:text-gold-bright">
|
||||
Na kávu
|
||||
</button>
|
||||
<button onClick={handleLogout} className="text-green-dim hover:text-gold">
|
||||
Odhlásiť
|
||||
</button>
|
||||
</div>
|
||||
<HeaderMenu
|
||||
username={account?.username}
|
||||
onHistory={() => navigate('/history')}
|
||||
onDonate={() => navigate('/donate')}
|
||||
onLogout={handleLogout}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div className="flex flex-col gap-3 mb-6">
|
||||
|
||||
@@ -2,6 +2,7 @@ import { useNavigate } from 'react-router-dom';
|
||||
import type { PlayerInfo } from '../types';
|
||||
import { computeTotal } from '../lib/standings';
|
||||
import { leaveGame } from '../lib/leaveGame';
|
||||
import { displayName } from '../lib/names';
|
||||
|
||||
interface Props {
|
||||
players: PlayerInfo[];
|
||||
@@ -30,7 +31,7 @@ export default function GameOver({ players, standings }: Props) {
|
||||
>
|
||||
<div className="flex items-center gap-3">
|
||||
<span className="text-2xl w-8">{medals[i]}</span>
|
||||
<span className="font-serif text-green-score">{p.name}</span>
|
||||
<span className="font-serif text-green-score">{displayName(p.name)}</span>
|
||||
</div>
|
||||
<span className={`font-serif text-xl ${i === 0 ? 'text-gold-bright' : 'text-gold'}`}>
|
||||
{p.total}
|
||||
|
||||
@@ -3,8 +3,9 @@ import { useNavigate } from 'react-router-dom';
|
||||
import { useGameStore } from '../store/gameStore';
|
||||
import { emit } from '../lib/socket';
|
||||
import { leaveGame } from '../lib/leaveGame';
|
||||
import { computePlayable } from '../lib/gameRules';
|
||||
import { computePlayable, stashWinner } from '../lib/gameRules';
|
||||
import { computeTotal } from '../lib/standings';
|
||||
import { displayName } from '../lib/names';
|
||||
import { useIsDesktop } from '../lib/useIsDesktop';
|
||||
import { useFitScale } from '../lib/useFitScale';
|
||||
import Hand from '../components/Hand';
|
||||
@@ -14,9 +15,14 @@ import Standings from '../components/Standings';
|
||||
import PlayerCircle from '../components/PlayerCircle';
|
||||
import FaceDownCards from '../components/FaceDownCards';
|
||||
import GameOver from './GameOver';
|
||||
import type { PlayerInfo, StashData } from '../types';
|
||||
import type { Hand as HandCards, PlayerInfo, StashData } from '../types';
|
||||
|
||||
const TRICK_LINGER_MS = 3000;
|
||||
// A completed trick stays face-up for SETTLE_MS (so the last card visibly joins
|
||||
// the pile), then is swept towards the winner over COLLECT_MS.
|
||||
const SETTLE_MS = 1100;
|
||||
const COLLECT_MS = 550;
|
||||
// Sweep direction by the winner's seat offset from me: 0=me(bottom) 1=left 2=top 3=right.
|
||||
const COLLECT_BY_OFFSET = ['collect-bottom', 'collect-left', 'collect-top', 'collect-right'];
|
||||
|
||||
export default function GameTable() {
|
||||
const navigate = useNavigate();
|
||||
@@ -28,28 +34,75 @@ export default function GameTable() {
|
||||
const gameStatus = useGameStore((s) => s.gameStatus);
|
||||
const hand = useGameStore((s) => s.hand);
|
||||
|
||||
// Hold the last completed trick visible for TRICK_LINGER_MS after it finishes.
|
||||
const [lingeredStash, setLingeredStash] = useState<StashData | null>(null);
|
||||
const lingerTimer = useRef<ReturnType<typeof setTimeout> | null>(null);
|
||||
// Once a completed trick has been swept away, its key is remembered here so it
|
||||
// is not shown again while we wait for the winner to lead the next trick.
|
||||
const [dismissedKey, setDismissedKey] = useState<string | null>(null);
|
||||
// Turns on for the collect (fly-to-winner) phase, after the settle pause.
|
||||
const [collecting, setCollecting] = useState(false);
|
||||
|
||||
const previousStash = gameStatus?.status.previous_stash ?? null;
|
||||
// Every game_status payload recreates the stash object, so identify the trick
|
||||
// by content — the timer must restart only when a *different* trick completes.
|
||||
// by content — the timers must restart only when a *different* trick completes.
|
||||
const previousStashKey = previousStash
|
||||
? `${previousStash.first_player}:${JSON.stringify(previousStash.cards)}`
|
||||
: null;
|
||||
|
||||
// On first load of an already-running game (reconnect / restore-on-startup) a
|
||||
// completed trick is already present; adopt it as "already swept" so we don't
|
||||
// replay a stale sweep over the live board. A freshly started game has no
|
||||
// completed trick at this point, so its very first trick still animates.
|
||||
const booted = useRef(false);
|
||||
useEffect(() => {
|
||||
if (!previousStash) return;
|
||||
setLingeredStash(previousStash);
|
||||
if (lingerTimer.current) clearTimeout(lingerTimer.current);
|
||||
lingerTimer.current = setTimeout(() => setLingeredStash(null), TRICK_LINGER_MS);
|
||||
if (booted.current || !gameStatus) return;
|
||||
booted.current = true;
|
||||
if (previousStashKey) setDismissedKey(previousStashKey);
|
||||
}, [gameStatus, previousStashKey]);
|
||||
|
||||
useEffect(() => {
|
||||
if (!previousStashKey) return;
|
||||
setCollecting(false);
|
||||
const settle = setTimeout(() => setCollecting(true), SETTLE_MS);
|
||||
const done = setTimeout(() => {
|
||||
setCollecting(false);
|
||||
setDismissedKey(previousStashKey);
|
||||
}, SETTLE_MS + COLLECT_MS);
|
||||
return () => {
|
||||
if (lingerTimer.current) clearTimeout(lingerTimer.current);
|
||||
clearTimeout(settle);
|
||||
clearTimeout(done);
|
||||
};
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [previousStashKey]);
|
||||
|
||||
// A just-completed trick that hasn't been swept away yet always wins the centre
|
||||
// — even once the winner has already led the next trick. The engine reveals that
|
||||
// lead card (and, at a round boundary, the next bidding phase) the instant the
|
||||
// 4th card lands, so without holding the pile here the sweep would be cut off
|
||||
// after every trick. Reading `previousStash` synchronously (rather than a state
|
||||
// set in an effect) also means the pile never blinks to empty on the frame the
|
||||
// 4th card lands — the last card simply joins the three already there, then the
|
||||
// whole pile is collected before the next trick takes over.
|
||||
const finishing = previousStashKey !== null && previousStashKey !== dismissedKey;
|
||||
|
||||
// The new round's dealt hand arrives (via `player_cards`) the instant the last
|
||||
// trick of the previous round is played, but the centre oval is still sweeping
|
||||
// that trick away — hold an empty hand on screen (the previous round's last card
|
||||
// really was just played, the engine just never broadcasts that "0 cards" beat on
|
||||
// its own since it deals the new round in the same step) until the sweep finishes,
|
||||
// so the new cards don't appear before the previous round has visibly wrapped up.
|
||||
// Mid-round plays (round number unchanged) still update instantly, since that's
|
||||
// just the player's own card leaving their hand, not a fresh deal.
|
||||
const [displayedHand, setDisplayedHand] = useState<HandCards>(hand);
|
||||
const lastAppliedRoundRef = useRef<number | null>(gameStatus?.round_number ?? null);
|
||||
useEffect(() => {
|
||||
if (!gameStatus) return;
|
||||
const isNewRound = gameStatus.round_number !== lastAppliedRoundRef.current;
|
||||
if (isNewRound && finishing) {
|
||||
setDisplayedHand({}); // last card of the previous round is gone; new deal waits
|
||||
return;
|
||||
}
|
||||
lastAppliedRoundRef.current = gameStatus.round_number;
|
||||
setDisplayedHand(hand);
|
||||
}, [hand, gameStatus, finishing]);
|
||||
|
||||
if (!gameStatus || !myPlayer) {
|
||||
return <p className="text-center text-green-dim pt-20 font-serif italic">Načítava sa…</p>;
|
||||
}
|
||||
@@ -73,14 +126,27 @@ export default function GameTable() {
|
||||
const myTurnToPlay = isPlayPhase && active_player === myOrder;
|
||||
|
||||
const activeCards = active_stash ? Object.keys(active_stash.cards).length : 0;
|
||||
const displayedStash: StashData | null =
|
||||
activeCards > 0 && active_stash ? active_stash : lingeredStash ?? null;
|
||||
const displayedStash: StashData | null = finishing
|
||||
? previousStash
|
||||
: activeCards > 0 && active_stash
|
||||
? active_stash
|
||||
: null;
|
||||
|
||||
const playableKeys = myTurnToPlay && active_stash
|
||||
? computePlayable(hand, active_stash.cards[String(active_stash.first_player)]?.color ?? null)
|
||||
// During the collect phase, sweep the pile towards whoever won it.
|
||||
const collectAnim =
|
||||
collecting && finishing && previousStash
|
||||
? COLLECT_BY_OFFSET[(stashWinner(previousStash) - myOrder + 4) % 4]
|
||||
: null;
|
||||
|
||||
// Block play until the previous trick's sweep animation has finished — otherwise
|
||||
// the winner could lead the next card while the pile is still visibly clearing.
|
||||
const canPlayNow = myTurnToPlay && !finishing;
|
||||
|
||||
const playableKeys = canPlayNow && active_stash
|
||||
? computePlayable(displayedHand, active_stash.cards[String(active_stash.first_player)]?.color ?? null)
|
||||
: undefined;
|
||||
|
||||
const activePlayerName = players.find((p) => p.order === active_player)?.name ?? '';
|
||||
const activePlayerName = displayName(players.find((p) => p.order === active_player)?.name);
|
||||
|
||||
// Seat mapping relative to "Ty": left / across / right.
|
||||
const seat = (offset: number): PlayerInfo | undefined =>
|
||||
@@ -140,7 +206,7 @@ export default function GameTable() {
|
||||
{opponents.map((p) => (
|
||||
<div key={p.order} className="text-center">
|
||||
<div className="uppercase tracking-[.1em] text-green-dim mb-0.5" style={{ fontSize: 11 }}>
|
||||
{p.name}
|
||||
{displayName(p.name)}
|
||||
</div>
|
||||
<div className="font-serif text-green-score leading-none" style={{ fontSize: compact ? 16 : 20 }}>
|
||||
{computeTotal(standings, p.order)}
|
||||
@@ -159,8 +225,14 @@ export default function GameTable() {
|
||||
);
|
||||
|
||||
// Center of the oval: trick during play, guess controls during bidding.
|
||||
const ovalContent = isPlayPhase ? (
|
||||
// `finishing` also keeps the trick on screen while the round's *last* stash is
|
||||
// swept away: the engine advances to the next round's bidding the instant the
|
||||
// 4th card lands, so `isPlayPhase` flips to false immediately — without this,
|
||||
// that final trick would vanish straight into the guess controls with no sweep.
|
||||
const ovalContent = isPlayPhase || finishing ? (
|
||||
<div style={collectAnim ? { animation: `${collectAnim} ${COLLECT_MS}ms ease-in both` } : undefined}>
|
||||
<Trick stash={displayedStash} players={players} myOrder={myOrder} />
|
||||
</div>
|
||||
) : (
|
||||
active_round_guesses !== undefined && active_player !== undefined ? (
|
||||
<GuessControls
|
||||
@@ -175,14 +247,14 @@ export default function GameTable() {
|
||||
|
||||
const topSeat = (
|
||||
<div className="flex flex-col items-center gap-1.5">
|
||||
<PlayerCircle name={topP?.name ?? '—'} {...seatProps(topP?.order)} size={desktop ? 64 : 52} />
|
||||
<PlayerCircle name={displayName(topP?.name) || '—'} {...seatProps(topP?.order)} size={desktop ? 64 : 52} />
|
||||
<FaceDownCards count={cardsInHandOf(topP?.order)} direction="row" desktop={desktop} />
|
||||
</div>
|
||||
);
|
||||
|
||||
const sideSeat = (p?: PlayerInfo) => (
|
||||
<div className="flex flex-col items-center gap-1.5">
|
||||
<PlayerCircle name={p?.name ?? '—'} {...seatProps(p?.order)} size={desktop ? 60 : 48} />
|
||||
<PlayerCircle name={displayName(p?.name) || '—'} {...seatProps(p?.order)} size={desktop ? 60 : 48} />
|
||||
<FaceDownCards count={cardsInHandOf(p?.order)} direction="col" desktop={desktop} />
|
||||
</div>
|
||||
);
|
||||
@@ -194,7 +266,7 @@ export default function GameTable() {
|
||||
);
|
||||
|
||||
const handArea = (
|
||||
<Hand hand={hand} myTurn={myTurnToPlay} isPlayPhase={isPlayPhase} playableKeys={playableKeys} desktop={desktop} />
|
||||
<Hand hand={displayedHand} myTurn={canPlayNow} isPlayPhase={isPlayPhase} playableKeys={playableKeys} desktop={desktop} />
|
||||
);
|
||||
|
||||
// ── DESKTOP LAYOUT ───────────────────────────────────────────────
|
||||
|
||||
@@ -3,6 +3,7 @@ import { useNavigate } from 'react-router-dom';
|
||||
import { useGameStore } from '../store/gameStore';
|
||||
import { emit, socket } from '../lib/socket';
|
||||
import { useIsDesktop } from '../lib/useIsDesktop';
|
||||
import { displayName } from '../lib/names';
|
||||
import type { GameDetail, GameDetailRound } from '../types';
|
||||
|
||||
function fmtDate(iso: string | null): string {
|
||||
@@ -62,7 +63,7 @@ export default function History() {
|
||||
className="flex-1 min-w-0 text-left px-4 py-3 hover:bg-white/[.02] transition-colors"
|
||||
>
|
||||
<p className="font-serif text-green-score truncate">{g.name || 'Hra'}</p>
|
||||
<p className="text-xs text-green-dim mt-1 truncate">{g.players.join(', ')}</p>
|
||||
<p className="text-xs text-green-dim mt-1 truncate">{g.players.map((n) => displayName(n)).join(', ')}</p>
|
||||
<p className="text-xs text-[#7a7058] mt-0.5">
|
||||
{fmtDate(g.created_at)} · {g.completed ? 'dohraná' : 'predčasne ukončená'}
|
||||
</p>
|
||||
@@ -128,7 +129,7 @@ function GameDetailView({ detail, onBack }: { detail: GameDetail; onBack: () =>
|
||||
return r.won ? (
|
||||
<span className="font-serif" style={{ fontSize: 14, color: '#c8bb95' }}>{r.points}</span>
|
||||
) : (
|
||||
<span className="font-serif line-through" style={{ fontSize: 14, color: '#7a6e4a' }}>{r.guess}</span>
|
||||
<span className="font-serif" style={{ fontSize: 14, color: '#7a6e4a' }}>{r.guess}</span>
|
||||
);
|
||||
};
|
||||
const seriesTotal = (s: number | undefined, seat: number) =>
|
||||
@@ -158,7 +159,7 @@ function GameDetailView({ detail, onBack }: { detail: GameDetail; onBack: () =>
|
||||
}}
|
||||
>
|
||||
<span className="uppercase text-gold" style={{ letterSpacing: '.06em', fontSize: 11 }}>
|
||||
{p.username}
|
||||
{displayName(p.username)}
|
||||
</span>
|
||||
<span className="font-serif text-gold-dim" style={{ fontWeight: 700, fontSize: 16 }}>
|
||||
{totals[c]}
|
||||
|
||||
@@ -3,6 +3,7 @@ import { useNavigate, useParams } from 'react-router-dom';
|
||||
import { useGameStore } from '../store/gameStore';
|
||||
import { emit } from '../lib/socket';
|
||||
import { leaveGame } from '../lib/leaveGame';
|
||||
import { displayName } from '../lib/names';
|
||||
|
||||
export default function Lobby() {
|
||||
const { gid } = useParams<{ gid: string }>();
|
||||
@@ -65,11 +66,35 @@ export default function Lobby() {
|
||||
return (
|
||||
<div key={order} className="flex items-center gap-3">
|
||||
<span className={`text-lg ${p ? 'text-gold' : 'text-[#7a7058]'}`}>
|
||||
{p ? '✦' : '○'}
|
||||
{p ? (p.is_bot ? '⚙' : '✦') : '○'}
|
||||
</span>
|
||||
<span className={p ? 'font-serif text-green-score' : 'text-green-dim italic'}>
|
||||
{p ? `${p.name}${myPlayer?.order === p.order ? ' (ty)' : ''}` : 'Čaká sa…'}
|
||||
{p ? `${displayName(p.name)}${myPlayer?.order === p.order ? ' (ty)' : ''}` : 'Čaká sa…'}
|
||||
</span>
|
||||
{isHost && p?.is_bot && (
|
||||
<button
|
||||
onClick={() => gid && emit.removeBot(gid, order)}
|
||||
className="ml-auto text-xs text-green-dim hover:text-gold"
|
||||
>
|
||||
Odobrať
|
||||
</button>
|
||||
)}
|
||||
{isHost && !p && (
|
||||
<span className="ml-auto flex gap-2">
|
||||
<button
|
||||
onClick={() => gid && emit.addBot(gid)}
|
||||
className="px-2 py-0.5 rounded-lg text-xs border border-gold/30 text-gold hover:bg-gold hover:text-table transition-colors"
|
||||
>
|
||||
+ Bot
|
||||
</button>
|
||||
<button
|
||||
onClick={() => gid && emit.addBot(gid, 'neural')}
|
||||
className="px-2 py-0.5 rounded-lg text-xs border border-gold/30 text-gold hover:bg-gold hover:text-table transition-colors"
|
||||
>
|
||||
+ AI bot
|
||||
</button>
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
|
||||
@@ -11,6 +11,7 @@ export interface PlayerInfo {
|
||||
name: string;
|
||||
connected: boolean;
|
||||
player_id?: number;
|
||||
is_bot?: boolean;
|
||||
}
|
||||
|
||||
export interface MyPlayer {
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# RL trening (rl/model.py, rl/selfplay.py, rl/train.py) -- zamerne oddelene
|
||||
# od requirements.txt: server ani Docker image torch nepotrebuju, boti v hre
|
||||
# pouzivaju len cisto-Python rl/players.py (a neskor natrenovane vahy cez
|
||||
# torch az ked sa neuralny bot nasadi).
|
||||
torch>=2.4
|
||||
+181
@@ -0,0 +1,181 @@
|
||||
# RL bot pre bridzik — navrh (2026-07-01)
|
||||
|
||||
Ciel: naucit sa principy self-play reinforcement learningu (v duchu AlphaGo Zero)
|
||||
na praktickom priklade — natrenovat sietovy policy pre hranie bridziku. Zamerne
|
||||
zjednodusene oproti AlphaGo Zero: bez MCTS (skryta informacia neumoznuje priamy
|
||||
prehladavaci strom), cisty self-play policy gradient (PPO/REINFORCE) nad `bridzik.py`
|
||||
enginom.
|
||||
|
||||
## Kluc: `Round` je nezavisla epizoda
|
||||
|
||||
Bodovanie (`Round.get_points_summary`) je cisto lokalne pre jedno kolo — nezavisi
|
||||
od predoslych ani nasledujucich kol, len od tipu a poctu kopiek v danom kole.
|
||||
Netreba teda simulovat cely `Bridzik`/`Series` state machine na trening — staci
|
||||
instanciovat `Round(round_number, first_player, shuffler)` priamo, opakovane,
|
||||
s roznymi `round_number` (0-7, teda 8 az 1 karta v ruke). Kazdy `Round` je
|
||||
samostatna self-play epizoda.
|
||||
|
||||
Vyhody:
|
||||
- jednoduchsi self-play loop (ziadne series/game bookkeeping)
|
||||
- vela nezavislych epizod, jednoduchá paralelizacia
|
||||
|
||||
**Curriculum vs. uniformne samplovanie (bod 4).** Povodny napad "zacat na
|
||||
`round_number=7`" je zavadzajuci: pri 1 karte je tip masked na {0,1} a jedina
|
||||
karta je vynutena — nula card-play rozhodnuti, ziadne ucenie, len smoke-test.
|
||||
Realne ucenie je v kolach `round_number` 0-3 (6-8 kariet). Preto:
|
||||
- `round_number` je aj tak v observacii, takze **default = uniformne samplovat
|
||||
`round_number` 0-7** a nechat siet zdielat vahy naprieč velkostami ruk.
|
||||
- Ak curriculum, tak **od tazkych (viac kariet) k lahsim**, nie naopak; alebo
|
||||
aspon uniformne s miernym zvyhodnenim tazsich kol.
|
||||
- `round_number=7` drzat len ako sanity/smoke test pipeline, nie ako trening.
|
||||
|
||||
## 1. State encoding (observation)
|
||||
|
||||
Spolocne pre guess aj play fazu, budovane z `Round` objektu pre daneho hraca:
|
||||
|
||||
- **vlastna ruka** — 32-dim multi-hot (4 farby x 8 hodnot)
|
||||
- **round_number** — one-hot (8) alebo normalizovane cislo (urcuje velkost ruk)
|
||||
- **tipy vsetkych 4 hracov** — 4x (flag "uz tipoval" + normalizovana hodnota)
|
||||
- **vlastny tip** (po tipnuti) — kriticke pre play fazu (viem, ci este potrebujem
|
||||
vyhrat kopku, alebo sa jej mam vyhybat)
|
||||
- **kolko kopiek uz kazdy hrac vyhral v tomto kole** — 4 scalars, odvodene
|
||||
z dokoncenych `Stash` objektov v `self.stashes`
|
||||
- **aktualna kopka v procese** — 4 sloty (karta alebo prazdne) + `first_player`
|
||||
aktualnej kopky
|
||||
- **uz odohrane/videne karty v tomto kole** — 32-dim multi-hot (bod 1). KRITICKE:
|
||||
bez toho observacia NIE JE Markovovska. Pri viac kartach (round_number 0-2)
|
||||
su dve rovnake ruky s rovnakou aktualnou kopkou, ale roznou historiou uz
|
||||
odohranych kariet, rozne stavy s roznym optimalnym tahom (vies, ci este visi
|
||||
eso/cerven). Bez tejto zlozky sa siet nemoze naucit card-counting a strop hry
|
||||
ostane nizky. Kodovat karty odohrane v predoslych dokoncenych `Stash`-och
|
||||
(mimo tvojej ruky a mimo aktualnej rozohranej kopky).
|
||||
- **(volitelne, neskor — bod 8) znama neúčasť supperov vo farbe (voids)** —
|
||||
ked supper neprizna vynasanu farbu, prezradi void → per-hrac x per-farba
|
||||
flag. Silna informacia, ale nechat na neskorsie rozsirenie.
|
||||
|
||||
**Egocentricka rotacia (bod 2) — povinny invariant.** Aby parameter sharing
|
||||
medzi 4 sedadlami fungoval, VSETKY 4-hracske vektory (tipy, pocty vyhranych
|
||||
kopiek, sloty aktualnej kopky, `first_player`) musia byt rotovane tak, ze
|
||||
"ja" = index 0 a ostatni relativne (+1, +2, +3 v smere hry). Toto zapisat do
|
||||
`encoding.py` ako tvrdy invariant + unit test — je to najpravdepodobnejsie
|
||||
miesto tichej chyby, ktora pokazi ucenie.
|
||||
|
||||
Zamerne vynechane: priebezne skore/standings naprieč hrou — kedze odmena je
|
||||
per-round nezavisla, optimalne rozhodnutie v danom kole na standings nezavisi.
|
||||
|
||||
**Velkost observacie (bod 5).** Povodny odhad ~80-100 floatov je podstrelený.
|
||||
Ak sa 4 sloty aktualnej kopky koduju one-hot (4x32=128) + ruka 32 + videne
|
||||
karty 32 + tipy/pocty/round_number/first_player, realny `input_dim` je skor
|
||||
~200. Nie je to problem, len podla toho nastavit vstupnu vrstvu siete.
|
||||
|
||||
## 2. Akcny priestor + maskovanie
|
||||
|
||||
- **Guess**: 9 kategorii (0-8), maskovane na `0..(8-round_number)`; pre 4.
|
||||
(posledneho) tipujuceho naviac zamaskovat hodnotu, ktora by sposobila
|
||||
`BridzikException` (sucet tipov = pocet kopiek) — vypocitatelne vopred
|
||||
z `self.guesses`. Zakazana hodnota = `(8-round_number) - sum(3 tipov)`;
|
||||
ak vyjde mimo `0..(8-round_number)`, je uz aj tak nelegalna a nemaskuje sa
|
||||
nic navyse (osetrit rozsah).
|
||||
- **Play card**: 32 kategorii (rovnaka indexacia farba+hodnota ako hand-encoding),
|
||||
maskovane na karty, ktore hrac realne ma A splnaju follow-suit pravidlo
|
||||
(rovnaka logika ako v `Round.play_card`: farba prvej karty v kopke, inak
|
||||
povinna cervena ak ju hrac ma).
|
||||
|
||||
## 3. Sieť
|
||||
|
||||
Zdielany "trup" (2 hidden layers, ~128-256 neuronov, ReLU) nad observation
|
||||
vektorom (~80-100 floatov), s troma vystupmi:
|
||||
|
||||
- guess head (9 logitov)
|
||||
- play head (32 logitov)
|
||||
- value head (1 scalar) — odhad ocakavanej odmeny do konca kola (baseline
|
||||
pre actor-critic)
|
||||
|
||||
`input_dim` nastavit podla realnej velkosti observacie (~200, viz bod 5
|
||||
v sekcii 1), nie podla povodneho ~80-100.
|
||||
|
||||
Fazovy flag v observacii + maskovanie urcuje, ktora hlava je pouzitelna
|
||||
v danom kroku (guess a play fazy sa nikdy neprelinaju).
|
||||
|
||||
## 4. Odmena a trening
|
||||
|
||||
- Odmena = 0 pocas kola; na konci kola kazdy hrac dostane
|
||||
`points_summary[player]` (0 alebo `10+guess`) ako terminalnu odmenu za
|
||||
VSETKY svoje rozhodnutia v danom kole (guess + vsetky `play_card` tahy).
|
||||
Sparse terminal reward, ziadne discountovanie netreba — `gamma=1` (bod 9),
|
||||
kolo ma max 9 rozhodnuti na hraca: round 0 = 1 tip + 8 kariet.
|
||||
- Algoritmus: self-play PPO (prip. najprv jednoduchsie REINFORCE + baseline),
|
||||
jedna zdielana siet hra vsetkych 4 hracov v kazdom `Round` (parameter
|
||||
sharing, rovnaky princip ako AlphaGo Zero).
|
||||
- **Normalizacia odmeny (bod 7).** `10+guess` je v rozsahu 10-18 a lisi sa
|
||||
per kolo; pri miesanych `round_number` to zvysuje varianciu policy gradientu.
|
||||
Standardizovat advantage per batch (odcitat priemer, delit std) — bezna
|
||||
PPO praktika, tu je nutnejsia kvoli rozne velkym odmenam.
|
||||
- Paralelizacia: `Round` instancie su nezavisle bez shared state, self-play
|
||||
generovanie sa da paralelizovat cez multiprocessing naprieč jadrami CPU
|
||||
(pripadne batchovanim viacerych epizod naraz cez sietovy forward).
|
||||
|
||||
**Caveat: hra nie je zero-sum (bod 3).** Je to 4-hracska general-sum hra —
|
||||
viacero hracov moze naraz trafit tip a vsetci skoruju, zaroven sa o kopky
|
||||
sutazi (`sum(kopky) = pocet kopiek`). Self-play so zdielanymi vahami preto
|
||||
NEMA konvergencne zaruky ako AlphaGo Zero (2-hracska zero-sum); skonverguje
|
||||
k *nejakemu* equilibriu, nie nutne k optimu, a moze oscilovat. Na ucebny
|
||||
ciel to staci, ale: (a) nepredavat si to ako "AlphaZero-grade optimalitu",
|
||||
(b) sledovat progres proti FIXNYM baseline-om (sekcia 5), nie len podla
|
||||
self-play reward, ktory sa hybe s protihracom.
|
||||
|
||||
## 5. Vyhodnotenie
|
||||
|
||||
- priemerne body/kolo oproti baseline (nahodny legalny hrac, Monte Carlo
|
||||
heuristicky tipper — pozri sekciu nizsie)
|
||||
- presnost tipu (% kôl, kde sa tip presne trafil) — interpretovatelnejsia
|
||||
metrika nez surove body
|
||||
|
||||
## Alternativa/doplnok pre tipovaciu fazu: Monte Carlo namiesto siete
|
||||
|
||||
Kedze tipovanie je v podstate odhad pravdepodobnosti pri neznamom rozdeleni
|
||||
zvysnych kariet, da sa riesit aj bez siete:
|
||||
|
||||
1. **Naivna MC simulacia** — vygenerovat vela nahodnych rozdeleni zvysnych
|
||||
kariet medzi ostatnych 3 hracov, odsimulovat kolo s jednoduchou heuristickou
|
||||
hracou strategiou, spocitat rozdelenie poctu vlastnych kopiek. POZOR: kedze
|
||||
bodujeme len presnu zhodu, spravny cieľ je **mod** rozdelenia, nie priemer.
|
||||
POZOR 2 (bod 6 — "discard pile"): `deal_starting_cards` zahodí prvych
|
||||
`4*round_number` kariet (`round_cards[4*round_number:]`), takze v kole NIE
|
||||
su rozdane vsetky karty. MC teda z `32 - vlastna_ruka` kariet rozdá kazdemu
|
||||
z 3 supperov len `(8-round_number)` kariet a **zvysok necha v neznamej kope
|
||||
mimo hru** — nerozdavat vsetko medzi supperov, inak nadhodnotis, kolko
|
||||
vysokych kariet/cervene supperi drzia.
|
||||
2. **Silnejsia verzia** — rovnaky MC rollout, ale simulovat zvysok kola
|
||||
s uz natrenovanou card-play sietou namiesto naivnej heuristiky (analogia
|
||||
MCTS + value network v AlphaZero namiesto ciste nahodnych rolloutov).
|
||||
|
||||
Toto sa da pouzit ako rychly heuristicky baseline bez trenovania siete na
|
||||
tipovanie vobec, alebo ako silnejsi hybrid s uz existujucou play sietou.
|
||||
|
||||
## Poradie implementacie
|
||||
|
||||
1. `rl/encoding.py` — cistě funkcie observation + mask (nad `Round` objektom),
|
||||
testovatelne izolovane. Uz tu zapracovat bod 1 (videne karty) a bod 2
|
||||
(egocentricka rotacia). Unit/property testy: maska NIKDY nepovoli tah, ktory
|
||||
`Round.play_card`/`add_player_guess` odmietne (fuzz-test proti enginu);
|
||||
rotacia je konzistentna pre vsetky 4 sedadla.
|
||||
2. `rl/env.py` — step/reset wrapper okolo jedneho `Round` (nie celeho `Bridzik`)
|
||||
3. **baseline hraci + evaluacny harness UZ TU** (nahodny legalny hrac, MC
|
||||
heuristicky tipper podla sekcie vyssie) — nech je metrika k dispozicii od
|
||||
prvej trenovacej epochy a da sa sledovat progres (bod 3: proti fixnym
|
||||
baseline-om, nie len self-play reward).
|
||||
4. sieť (PyTorch, trup + 3 hlavy) — `input_dim` podla realnej velkosti obs (~200)
|
||||
5. self-play generator (paralelne `Round` epizody)
|
||||
6. PPO update krok (advantage standardizovat per batch — bod 7)
|
||||
7. (neskor, volitelne) rozsirit na cely `Series`/`Bridzik` self-play, ak by
|
||||
sa ukazalo, ze cross-round dynamika (rotacia first_player a pod.) predsa
|
||||
len nieco mení — podla bodu 1 to nie je ocakavane
|
||||
|
||||
## Technologie
|
||||
|
||||
- **PyTorch** — samostatny `requirements-rl.txt`, oddeleny od zakladneho
|
||||
`requirements.txt` projektu
|
||||
- vlastna mensia implementacia PPO (v duchu CleanRL) namiesto Stable-Baselines3/
|
||||
RLlib — nas pripad (multi-agent self-play so zdielanymi vahami, maskovanie
|
||||
akcii) sa bije s ich single-agent Gym abstrakciou viac, nez by pomohla
|
||||
+168
@@ -0,0 +1,168 @@
|
||||
"""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
|
||||
@@ -0,0 +1,70 @@
|
||||
"""Self-play prostredie nad jednym `Round`-om (viz rl/DESIGN.md).
|
||||
|
||||
Jedno kolo = jedna epizoda. Prostredie je multi-agentne a tahove: v kazdom
|
||||
kroku je na tahu prave jeden hrac (`Decision.player`), akciu zan doda
|
||||
volajuci (zdielana siet, heuristika, ...). Odmena je sparse a terminalna --
|
||||
`Round.get_points_summary()` pre vsetkych 4 hracov naraz na konci kola.
|
||||
|
||||
Akcie: v tipovacej faze index tipu 0..8, v hracej faze index karty 0..31
|
||||
(kanonicke cislovanie z rl/encoding.py). Legalne akcie urcuje
|
||||
`Decision.mask`; nelegalna akcia prebuble ako BridzikException z enginu.
|
||||
"""
|
||||
|
||||
from random import Random
|
||||
from typing import NamedTuple
|
||||
|
||||
from bridzik import Round, ROUNDS_PER_SERIES
|
||||
from rl.encoding import encode_observation, guess_mask, index_card, play_mask
|
||||
|
||||
PHASE_GUESS = 'guess'
|
||||
PHASE_PLAY = 'play'
|
||||
|
||||
|
||||
class Decision(NamedTuple):
|
||||
"""Jeden rozhodovaci bod: kto je na tahu, v akej faze, co vidi a co smie."""
|
||||
player: int
|
||||
phase: str
|
||||
obs: list
|
||||
mask: list
|
||||
|
||||
|
||||
class RoundEnv:
|
||||
def __init__(self, rng: Random = None):
|
||||
self.rng = rng if rng is not None else Random()
|
||||
self.round = None
|
||||
|
||||
def reset(self, round_number: int = None, first_player: int = None) -> Decision:
|
||||
"""Zacne novu epizodu; nezadane parametre sa sampluju uniformne."""
|
||||
if round_number is None:
|
||||
round_number = self.rng.randrange(ROUNDS_PER_SERIES)
|
||||
if first_player is None:
|
||||
first_player = self.rng.randrange(4)
|
||||
self.round = Round(round_number, first_player, shuffler=self.rng.shuffle)
|
||||
return self._decision()
|
||||
|
||||
def step(self, action: int):
|
||||
"""Vykona akciu hraca na tahu.
|
||||
|
||||
Vracia (decision, rewards, done): pocas kola (Decision, None, False),
|
||||
na konci kola (None, [body 4 hracov], True).
|
||||
"""
|
||||
if self.round is None or self.round.is_completed():
|
||||
raise RuntimeError('Epizoda nebezi -- najprv zavolaj reset().')
|
||||
player = self.round.get_active_player()
|
||||
if not self.round.is_guessing_completed():
|
||||
self.round.add_player_guess(player, action)
|
||||
else:
|
||||
self.round.play_card(player, index_card(action))
|
||||
if self.round.is_completed():
|
||||
return None, self.round.get_points_summary(), True
|
||||
return self._decision(), None, False
|
||||
|
||||
def _decision(self) -> Decision:
|
||||
player = self.round.get_active_player()
|
||||
if not self.round.is_guessing_completed():
|
||||
return Decision(player, PHASE_GUESS,
|
||||
encode_observation(self.round, player),
|
||||
guess_mask(self.round))
|
||||
return Decision(player, PHASE_PLAY,
|
||||
encode_observation(self.round, player),
|
||||
play_mask(self.round, player))
|
||||
@@ -0,0 +1,89 @@
|
||||
"""Evaluacny harness: odohra N kol medzi 4 hracmi a spocita metriky.
|
||||
|
||||
Metriky per hrac (viz rl/DESIGN.md, sekcia 5): priemerne body na kolo
|
||||
a presnost tipu (% kol s presne trafenym tipom). Sedadla sa medzi kolami
|
||||
rotuju, aby ziadny hrac nebol systematicky zvyhodneny poradim tipovania.
|
||||
|
||||
Spustenie ako skript porovna baseline botov:
|
||||
py -m rl.evaluate --rounds 500 --seed 7
|
||||
"""
|
||||
|
||||
import argparse
|
||||
from random import Random
|
||||
|
||||
from bridzik import ROUNDS_PER_SERIES
|
||||
from rl.env import PHASE_GUESS, RoundEnv
|
||||
|
||||
|
||||
def play_round(players: list, env: RoundEnv, round_number: int = None,
|
||||
first_player: int = None) -> list:
|
||||
"""Odohra jedno kolo; `players[seat]` rozhoduje za sedadlo `seat`.
|
||||
|
||||
Vrati body 4 sedadiel (`Round.get_points_summary()`).
|
||||
"""
|
||||
decision = env.reset(round_number, first_player)
|
||||
while True:
|
||||
seat = decision.player
|
||||
if decision.phase == PHASE_GUESS:
|
||||
action = players[seat].guess(env.round, seat)
|
||||
else:
|
||||
action = players[seat].play(env.round, seat)
|
||||
decision, rewards, done = env.step(action)
|
||||
if done:
|
||||
return rewards
|
||||
|
||||
|
||||
def evaluate(players: list, n_rounds: int, rng: Random = None,
|
||||
round_numbers: list = None) -> list:
|
||||
"""Odohra `n_rounds` kol s rotaciou sedadiel; vrati stats per hrac.
|
||||
|
||||
Vystup: zoznam dictov v poradi `players` --
|
||||
{'avg_points': float, 'hit_rate': float, 'rounds': int}.
|
||||
"""
|
||||
rng = rng if rng is not None else Random()
|
||||
env = RoundEnv(rng)
|
||||
points = [0] * 4
|
||||
hits = [0] * 4
|
||||
for i in range(n_rounds):
|
||||
round_number = rng.choice(round_numbers) if round_numbers \
|
||||
else rng.randrange(ROUNDS_PER_SERIES)
|
||||
# rotacia: sedadlo s obsadzuje players[(s + i) % 4]
|
||||
seating = [players[(s + i) % 4] for s in range(4)]
|
||||
rewards = play_round(seating, env, round_number)
|
||||
for seat in range(4):
|
||||
player_index = (seat + i) % 4
|
||||
points[player_index] += rewards[seat]
|
||||
hits[player_index] += rewards[seat] > 0
|
||||
return [{'avg_points': points[p] / n_rounds,
|
||||
'hit_rate': hits[p] / n_rounds,
|
||||
'rounds': n_rounds} for p in range(len(players))]
|
||||
|
||||
|
||||
def main():
|
||||
from rl.players import HeuristicPlayer, RandomPlayer
|
||||
|
||||
parser = argparse.ArgumentParser(description='Evaluacia baseline botov')
|
||||
parser.add_argument('--rounds', type=int, default=500)
|
||||
parser.add_argument('--seed', type=int, default=7)
|
||||
parser.add_argument('--mc-samples', type=int, default=100)
|
||||
args = parser.parse_args()
|
||||
|
||||
rng = Random(args.seed)
|
||||
lineups = [
|
||||
('4x random', [RandomPlayer(rng) for _ in range(4)]),
|
||||
('1x heuristika + 3x random',
|
||||
[HeuristicPlayer(rng, n_samples=args.mc_samples)]
|
||||
+ [RandomPlayer(rng) for _ in range(3)]),
|
||||
('4x heuristika',
|
||||
[HeuristicPlayer(rng, n_samples=args.mc_samples) for _ in range(4)]),
|
||||
]
|
||||
for label, players in lineups:
|
||||
stats = evaluate(players, args.rounds, rng)
|
||||
print(f'\n{label} ({args.rounds} kol):')
|
||||
for i, s in enumerate(stats):
|
||||
print(f' hrac {i}: {s["avg_points"]:6.2f} bodov/kolo, '
|
||||
f'tip trafeny {100 * s["hit_rate"]:5.1f} %')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -0,0 +1,69 @@
|
||||
"""Export vah natrenovanej siete do formatu pre cisto-Python inferenciu.
|
||||
|
||||
Torch je len trenovacia zavislost (host); produkcia hra cez rl/pure_net.py,
|
||||
ktory cita tento subor bez torch/numpy. Vahy sa uladaju ako base64 float32
|
||||
little-endian (bit-exact kopia checkpointu, ziadna strata presnosti).
|
||||
|
||||
Pouzitie:
|
||||
py -m rl.export rl/checkpoints/latest.pt rl/weights/neural-bot.json
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
|
||||
from rl.encoding import OBS_DIM
|
||||
from rl.train import load_checkpoint
|
||||
|
||||
DEFAULT_WEIGHTS_PATH = os.path.join('rl', 'weights', 'neural-bot.json')
|
||||
|
||||
|
||||
def _pack(tensor) -> dict:
|
||||
"""Tensor -> {shape, base64 float32 LE}. Cez struct, bez numpy -- tolist()
|
||||
vracia presne hodnoty float32, takze zapis je bit-exact."""
|
||||
import struct
|
||||
data = tensor.detach().to('cpu').contiguous().float()
|
||||
flat = data.reshape(-1).tolist()
|
||||
return {
|
||||
'shape': list(data.shape),
|
||||
'data': base64.b64encode(struct.pack(f'<{len(flat)}f', *flat)).decode('ascii'),
|
||||
}
|
||||
|
||||
|
||||
def export(checkpoint_path: str, out_path: str) -> dict:
|
||||
net = load_checkpoint(checkpoint_path)
|
||||
hidden = net.trunk[0].out_features
|
||||
payload = {
|
||||
'obs_dim': OBS_DIM,
|
||||
'hidden': hidden,
|
||||
'weights': {
|
||||
'trunk0_w': _pack(net.trunk[0].weight),
|
||||
'trunk0_b': _pack(net.trunk[0].bias),
|
||||
'trunk2_w': _pack(net.trunk[2].weight),
|
||||
'trunk2_b': _pack(net.trunk[2].bias),
|
||||
'guess_w': _pack(net.guess_head.weight),
|
||||
'guess_b': _pack(net.guess_head.bias),
|
||||
'play_w': _pack(net.play_head.weight),
|
||||
'play_b': _pack(net.play_head.bias),
|
||||
},
|
||||
}
|
||||
os.makedirs(os.path.dirname(out_path), exist_ok=True)
|
||||
with open(out_path, 'w') as f:
|
||||
json.dump(payload, f)
|
||||
return payload
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='Export vah pre pure-Python inferenciu')
|
||||
parser.add_argument('checkpoint', nargs='?', default='rl/checkpoints/latest.pt')
|
||||
parser.add_argument('out', nargs='?', default=DEFAULT_WEIGHTS_PATH)
|
||||
args = parser.parse_args()
|
||||
payload = export(args.checkpoint, args.out)
|
||||
size = os.path.getsize(args.out)
|
||||
print(f'Exportovane: {args.checkpoint} (hidden={payload["hidden"]}) '
|
||||
f'-> {args.out} ({size / 1024:.0f} kB)')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
+49
@@ -0,0 +1,49 @@
|
||||
"""Siet pre self-play PPO (viz rl/DESIGN.md, sekcia 3).
|
||||
|
||||
Zdielany trup nad observaciou z rl/encoding.py a tri hlavy:
|
||||
guess (9 logitov), play (32 logitov), value (1 skalar -- baseline pre
|
||||
actor-critic). Ktora policy hlava plati, urcuje faza rozhodnutia; nelegalne
|
||||
akcie sa odrezavaju maskou (logit -inf), takze distribucia nikdy nenavzorkuje
|
||||
tah, ktory by engine odmietol.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from rl.encoding import N_GUESS_ACTIONS, N_PLAY_ACTIONS, OBS_DIM
|
||||
|
||||
|
||||
class BridzikNet(nn.Module):
|
||||
def __init__(self, hidden: int = 256):
|
||||
super().__init__()
|
||||
self.trunk = nn.Sequential(
|
||||
nn.Linear(OBS_DIM, hidden), nn.ReLU(),
|
||||
nn.Linear(hidden, hidden), nn.ReLU(),
|
||||
)
|
||||
self.guess_head = nn.Linear(hidden, N_GUESS_ACTIONS)
|
||||
self.play_head = nn.Linear(hidden, N_PLAY_ACTIONS)
|
||||
self.value_head = nn.Linear(hidden, 1)
|
||||
|
||||
def forward(self, obs: torch.Tensor):
|
||||
"""obs (B, OBS_DIM) -> (guess_logits (B,9), play_logits (B,32), value (B,))."""
|
||||
h = self.trunk(obs)
|
||||
return self.guess_head(h), self.play_head(h), self.value_head(h).squeeze(-1)
|
||||
|
||||
|
||||
def masked_categorical(logits: torch.Tensor, mask: torch.Tensor) -> torch.distributions.Categorical:
|
||||
"""Kategoricka distribucia s nelegalnymi akciami odrezanymi na -inf.
|
||||
|
||||
`mask` je bool tensor rovnakeho tvaru ako `logits`; kazdy riadok musi mat
|
||||
aspon jednu povolenu akciu (garantuju masky z rl/encoding.py).
|
||||
"""
|
||||
return torch.distributions.Categorical(
|
||||
logits=logits.masked_fill(~mask, float('-inf'))
|
||||
)
|
||||
|
||||
|
||||
def obs_tensor(obs: list) -> torch.Tensor:
|
||||
return torch.tensor(obs, dtype=torch.float32)
|
||||
|
||||
|
||||
def mask_tensor(mask: list) -> torch.Tensor:
|
||||
return torch.tensor(mask, dtype=torch.bool)
|
||||
+264
@@ -0,0 +1,264 @@
|
||||
"""Baseline hraci pre evaluaciu a neskorsi warm-start siete (viz rl/DESIGN.md).
|
||||
|
||||
Spolocne rozhranie: `guess(rnd, seat) -> int` (tip 0..8) a
|
||||
`play(rnd, seat) -> int` (index karty 0..31). Hrac vidi len to, co by videl
|
||||
pri stole -- vlastnu ruku, tipy, dokoncene kopky a rozohranu kopku; do cudzich
|
||||
ruk nesiaha.
|
||||
"""
|
||||
|
||||
from collections import Counter
|
||||
from random import Random
|
||||
|
||||
from bridzik import cards, Card_colors, Stash
|
||||
from rl.encoding import (
|
||||
N_GUESS_ACTIONS, N_PLAY_ACTIONS,
|
||||
card_index, deduce_voids, guess_mask, legal_cards, play_mask,
|
||||
)
|
||||
|
||||
|
||||
class RandomPlayer:
|
||||
"""Uniformne nahodny legalny tah -- najslabsi mozny baseline."""
|
||||
|
||||
def __init__(self, rng: Random = None):
|
||||
self.rng = rng if rng is not None else Random()
|
||||
|
||||
def guess(self, rnd, seat: int) -> int:
|
||||
mask = guess_mask(rnd)
|
||||
return self.rng.choice([g for g in range(N_GUESS_ACTIONS) if mask[g]])
|
||||
|
||||
def play(self, rnd, seat: int) -> int:
|
||||
mask = play_mask(rnd, seat)
|
||||
return self.rng.choice([i for i in range(N_PLAY_ACTIONS) if mask[i]])
|
||||
|
||||
|
||||
def _strength(card) -> tuple:
|
||||
"""Absolutna sila karty: kazda cervena (tromf) bije kazdu necervenu."""
|
||||
return (card.color == Card_colors['HEARTS'], card.value.value)
|
||||
|
||||
|
||||
def _current_best(stash):
|
||||
"""Zatial vitazna karta rozohranej kopky (None ak sa este nevynieslo)."""
|
||||
first = stash.get_first_card() if stash is not None else None
|
||||
if first is None:
|
||||
return None
|
||||
best = first
|
||||
for card in stash.get_cards().values():
|
||||
if _beats(card, best):
|
||||
best = card
|
||||
return best
|
||||
|
||||
|
||||
def _beats(card, best) -> bool:
|
||||
"""Ci `card` prebije `best` (karta drziaca kopku; jej farba je smerodajna)."""
|
||||
if card.color == best.color:
|
||||
return card.value > best.value
|
||||
return card.color == Card_colors['HEARTS']
|
||||
|
||||
|
||||
def simulate_tricks(hands: dict, leader: int, rng: Random) -> list:
|
||||
"""Dohra kopky s nahodnou legalnou strategiou; vrati pocty vyhier hracov.
|
||||
|
||||
`hands` je dict seat -> zoznam kariet (rovnako velke ruky); zoznamy sa
|
||||
spotrebuju. Vitaza kopky urcuje enginovy Stash.get_winner() -- pravidla
|
||||
sa tu neduplikuju.
|
||||
"""
|
||||
tricks = [0] * 4
|
||||
for _ in range(len(hands[leader])):
|
||||
stash = Stash(leader)
|
||||
for i in range(4):
|
||||
seat = (leader + i) % 4
|
||||
card = rng.choice(legal_cards(hands[seat], stash.get_first_card()))
|
||||
hands[seat].remove(card)
|
||||
stash.add_card(seat, card)
|
||||
leader = stash.get_winner()
|
||||
tricks[leader] += 1
|
||||
return tricks
|
||||
|
||||
|
||||
def deal_consistent(unknown: list, hand_sizes: dict, voids: dict,
|
||||
rng: Random, max_tries: int = 20) -> dict:
|
||||
"""Nahodne rozdanie neznamych kariet superom respektujuce voidy.
|
||||
|
||||
Greedy priradenie po zamiesani (najviac obmedzeni hraci prvi); ak sa
|
||||
konzistentne rozdanie nepodari za `max_tries`, padne na rozdanie bez
|
||||
voidov (zriedkave, radsej mierne skreslena vzorka nez ziadna).
|
||||
Zvysok kariet ostava v odlozenej kope mimo hry.
|
||||
"""
|
||||
seats = sorted(hand_sizes, key=lambda s: len(voids.get(s, ())), reverse=True)
|
||||
pool = list(unknown)
|
||||
for _ in range(max_tries):
|
||||
rng.shuffle(pool)
|
||||
remaining = list(pool)
|
||||
hands = {}
|
||||
for seat in seats:
|
||||
hand, rest, banned = [], [], voids.get(seat, set())
|
||||
for card in remaining:
|
||||
if len(hand) < hand_sizes[seat] and card.color not in banned:
|
||||
hand.append(card)
|
||||
else:
|
||||
rest.append(card)
|
||||
if len(hand) < hand_sizes[seat]:
|
||||
break
|
||||
hands[seat] = hand
|
||||
remaining = rest
|
||||
else:
|
||||
return hands
|
||||
rng.shuffle(pool)
|
||||
idx = 0
|
||||
hands = {}
|
||||
for seat in seats:
|
||||
hands[seat] = pool[idx:idx + hand_sizes[seat]]
|
||||
idx += hand_sizes[seat]
|
||||
return hands
|
||||
|
||||
|
||||
def mc_guess_distribution(rnd, seat: int, n_samples: int, rng: Random) -> Counter:
|
||||
"""Monte Carlo odhad rozdelenia poctu vlastnych kopiek v kole.
|
||||
|
||||
Nezname karty sa v kazdej vzorke nahodne rozdelia ostatnym trom hracom
|
||||
-- kazdemu len (8 - round_number) kariet, zvysok ostava v odlozenej kope
|
||||
mimo hry (pozri DESIGN.md, pasca "discard pile"). Leader prvej kopky je
|
||||
v case tipovania neznamy (najvyssi tip), sampluje sa uniformne.
|
||||
"""
|
||||
hand = rnd.player_cards[seat]
|
||||
hand_size = 8 - rnd.round_number
|
||||
unknown = [c for c in cards if c not in hand]
|
||||
counts = Counter()
|
||||
for _ in range(n_samples):
|
||||
rng.shuffle(unknown)
|
||||
sim_hands = {seat: list(hand)}
|
||||
others = [s for s in range(4) if s != seat]
|
||||
for i, other in enumerate(others):
|
||||
sim_hands[other] = unknown[i * hand_size:(i + 1) * hand_size]
|
||||
tricks = simulate_tricks(sim_hands, rng.randrange(4), rng)
|
||||
counts[tricks[seat]] += 1
|
||||
return counts
|
||||
|
||||
|
||||
def finish_round(hands: dict, current_cards: dict, first_player: int,
|
||||
me: int, my_card, tricks: list, rng: Random) -> list:
|
||||
"""Dohra kolo od mojho tahu: dokonci rozohranu kopku (ja hram `my_card`,
|
||||
dalsi nahodne legalne) a zvysne kopky dohra nahodnou legalnou strategiou.
|
||||
`tricks` su uz vyhrane kopky (mutuje sa kopia volajuceho); vrati final."""
|
||||
stash = Stash(first_player)
|
||||
for seat, card in current_cards.items():
|
||||
stash.add_card(seat, card)
|
||||
stash.add_card(me, my_card)
|
||||
while not stash.is_completed():
|
||||
seat = stash.get_active_player()
|
||||
card = rng.choice(legal_cards(hands[seat], stash.get_first_card()))
|
||||
hands[seat].remove(card)
|
||||
stash.add_card(seat, card)
|
||||
leader = stash.get_winner()
|
||||
tricks[leader] += 1
|
||||
while hands[leader]:
|
||||
stash = Stash(leader)
|
||||
for i in range(4):
|
||||
seat = (leader + i) % 4
|
||||
card = rng.choice(legal_cards(hands[seat], stash.get_first_card()))
|
||||
hands[seat].remove(card)
|
||||
stash.add_card(seat, card)
|
||||
leader = stash.get_winner()
|
||||
tricks[leader] += 1
|
||||
return tricks
|
||||
|
||||
|
||||
class HeuristicPlayer:
|
||||
"""MC tipper + jednoducha hracia heuristika riadena vlastnym tipom."""
|
||||
|
||||
def __init__(self, rng: Random = None, n_samples: int = 100):
|
||||
self.rng = rng if rng is not None else Random()
|
||||
self.n_samples = n_samples
|
||||
|
||||
def guess(self, rnd, seat: int) -> int:
|
||||
counts = mc_guess_distribution(rnd, seat, self.n_samples, self.rng)
|
||||
mask = guess_mask(rnd)
|
||||
# najcastejsi LEGALNY pocet kopiek (mod rozdelenia, nie priemer --
|
||||
# boduje sa len presna zhoda); pri nule vzoriek pre legalny tip
|
||||
# rozhodne blizkost k celkovemu modu
|
||||
mode = counts.most_common(1)[0][0]
|
||||
legal = [g for g in range(N_GUESS_ACTIONS) if mask[g]]
|
||||
return max(legal, key=lambda g: (counts[g], -abs(g - mode)))
|
||||
|
||||
def play(self, rnd, seat: int) -> int:
|
||||
hand = rnd.player_cards[seat]
|
||||
stash = rnd.get_last_stash()
|
||||
allowed = legal_cards(hand, stash.get_first_card() if stash else None)
|
||||
need = rnd.guesses[seat] - rnd.get_stashes_winner_summary()[seat]
|
||||
best = _current_best(stash)
|
||||
|
||||
if best is None:
|
||||
# vynasam: chcem kopku -> najsilnejsia karta; nechcem -> najslabsia
|
||||
chosen = max(allowed, key=_strength) if need > 0 else min(allowed, key=_strength)
|
||||
else:
|
||||
winning = [c for c in allowed if _beats(c, best)]
|
||||
if need > 0 and winning:
|
||||
# ber kopku co najlacnejsie
|
||||
chosen = min(winning, key=_strength)
|
||||
elif need <= 0 and len(winning) < len(allowed):
|
||||
# kopku nechcem: zbav sa najsilnejsej neberucej karty
|
||||
chosen = max((c for c in allowed if not _beats(c, best)), key=_strength)
|
||||
elif need <= 0:
|
||||
# vsetko berie -> ber co najlacnejsie (setri silne karty netreba,
|
||||
# ale nizka karta drzi sancu, ze ma este niekto prebije)
|
||||
chosen = min(allowed, key=_strength)
|
||||
else:
|
||||
# kopku chcem, ale nic neberie -> odhod najslabsiu
|
||||
chosen = min(allowed, key=_strength)
|
||||
return card_index(chosen)
|
||||
|
||||
|
||||
class McPlayer(HeuristicPlayer):
|
||||
"""Heuristika s MC hracou fazou: kazdy kandidatsky tah sa ohodnoti
|
||||
simulaciami zvysku kola nad rozdaniami neznamych kariet konzistentnymi
|
||||
s dedukovanymi voidmi (`use_voids=False` = ablacia bez dedukcie).
|
||||
Tipovanie ostava MC tipper z HeuristicPlayer (pred prvou kartou niet
|
||||
z coho voidy dedukovat)."""
|
||||
|
||||
def __init__(self, rng: Random = None, n_samples: int = 100,
|
||||
play_samples: int = 24, use_voids: bool = True):
|
||||
super().__init__(rng, n_samples)
|
||||
self.play_samples = play_samples
|
||||
self.use_voids = use_voids
|
||||
|
||||
def play(self, rnd, seat: int) -> int:
|
||||
hand = rnd.player_cards[seat]
|
||||
stash = rnd.get_last_stash()
|
||||
first_card = stash.get_first_card() if stash else None
|
||||
candidates = legal_cards(hand, first_card)
|
||||
if len(candidates) == 1:
|
||||
return card_index(candidates[0])
|
||||
|
||||
target = rnd.guesses[seat]
|
||||
base_tricks = rnd.get_stashes_winner_summary()
|
||||
voids = deduce_voids(rnd) if self.use_voids else {}
|
||||
seen = set()
|
||||
played_count = Counter()
|
||||
for st in rnd.stashes:
|
||||
for other, card in st.get_cards().items():
|
||||
seen.add(card)
|
||||
played_count[other] += 1
|
||||
unknown = [c for c in cards if c not in seen and c not in hand]
|
||||
hand_size0 = 8 - rnd.round_number
|
||||
hand_sizes = {s: hand_size0 - played_count[s] for s in range(4) if s != seat}
|
||||
current_cards = stash.get_cards()
|
||||
|
||||
# spolocne rozdanie pre vsetkych kandidatov (common random numbers --
|
||||
# porovnavame tahy na tych istych svetoch, mensia variancia)
|
||||
scores = {card_index(c): 0 for c in candidates}
|
||||
for _ in range(self.play_samples):
|
||||
world = deal_consistent(unknown, hand_sizes, voids, self.rng)
|
||||
for candidate in candidates:
|
||||
sim_hands = {s: list(h) for s, h in world.items()}
|
||||
sim_hands[seat] = [c for c in hand if c != candidate]
|
||||
tricks = finish_round(sim_hands, dict(current_cards),
|
||||
stash.first_player, seat, candidate,
|
||||
list(base_tricks), self.rng)
|
||||
if tricks[seat] == target:
|
||||
scores[card_index(candidate)] += 1
|
||||
|
||||
best_index = max(scores, key=scores.get)
|
||||
if scores[best_index] == 0:
|
||||
# tip uz je (takmer) nedosiahnutelny -> aspon rozumny pravidlovy tah
|
||||
return super().play(rnd, seat)
|
||||
return best_index
|
||||
@@ -0,0 +1,42 @@
|
||||
"""Natrenovana siet ako hrac so standardnym rozhranim guess/play.
|
||||
|
||||
Rovnake rozhranie ako rl/players.py, takze funguje v rl/evaluate.py aj ako
|
||||
boti "mozog" v api/bots.py. Hrac vidi len observaciu + masku z rl/encoding.py
|
||||
-- z principu nemoze podvadzat (do cudzich ruk sa nedostane).
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from rl.encoding import encode_observation, guess_mask, play_mask
|
||||
from rl.model import BridzikNet, mask_tensor, masked_categorical, obs_tensor
|
||||
|
||||
|
||||
class NeuralPlayer:
|
||||
def __init__(self, net: BridzikNet, greedy: bool = True):
|
||||
self.net = net
|
||||
self.greedy = greedy # argmax pri evaluacii; sampling pre pestrost
|
||||
|
||||
def _act(self, rnd, seat: int, use_play_head: bool) -> int:
|
||||
obs = obs_tensor(encode_observation(rnd, seat)).unsqueeze(0)
|
||||
mask = mask_tensor(
|
||||
play_mask(rnd, seat) if use_play_head else guess_mask(rnd)
|
||||
).unsqueeze(0)
|
||||
self.net.eval()
|
||||
with torch.no_grad():
|
||||
guess_logits, play_logits, _ = self.net(obs)
|
||||
logits = play_logits if use_play_head else guess_logits
|
||||
logits = logits.masked_fill(~mask, float('-inf'))
|
||||
if self.greedy:
|
||||
return int(logits.argmax(dim=-1).item())
|
||||
return int(masked_categorical(logits, mask).sample().item())
|
||||
|
||||
def guess(self, rnd, seat: int) -> int:
|
||||
return self._act(rnd, seat, use_play_head=False)
|
||||
|
||||
def play(self, rnd, seat: int) -> int:
|
||||
return self._act(rnd, seat, use_play_head=True)
|
||||
|
||||
|
||||
def load_player(checkpoint_path: str, greedy: bool = True) -> NeuralPlayer:
|
||||
from rl.train import load_checkpoint
|
||||
return NeuralPlayer(load_checkpoint(checkpoint_path), greedy=greedy)
|
||||
+114
@@ -0,0 +1,114 @@
|
||||
"""Cisto-Python inferencia natrenovanej siete (stdlib only, bez torch/numpy).
|
||||
|
||||
Nacita vahy z exportu rl/export.py a implementuje forward pass MLP
|
||||
(trunk 2x ReLU + guess/play hlavy). Sluzi produkcnym botom v api/bots.py --
|
||||
torch ostava len trenovacia zavislost na hoste. Presnost overuje
|
||||
tests/test_pure_net.py porovnanim s torch vystupmi na zivych observaciach.
|
||||
|
||||
Vykon: ~250k nasobeni na tah (~desiatky ms) -- pri pauze medzi tahmi bota
|
||||
(BOT_MOVE_DELAY_SECONDS) nepostrehnutelne.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import struct
|
||||
|
||||
from rl.encoding import (
|
||||
N_GUESS_ACTIONS, N_PLAY_ACTIONS, OBS_DIM,
|
||||
encode_observation, guess_mask, play_mask,
|
||||
)
|
||||
|
||||
DEFAULT_WEIGHTS_PATH = os.path.join(
|
||||
os.path.dirname(__file__), 'weights', 'neural-bot.json'
|
||||
)
|
||||
|
||||
|
||||
def _unpack(entry: dict):
|
||||
"""{shape, base64 f32 LE} -> matica (list riadkov) alebo vektor."""
|
||||
flat = list(struct.unpack(
|
||||
f'<{_numel(entry["shape"])}f', base64.b64decode(entry['data'])
|
||||
))
|
||||
shape = entry['shape']
|
||||
if len(shape) == 1:
|
||||
return flat
|
||||
rows, cols = shape
|
||||
return [flat[r * cols:(r + 1) * cols] for r in range(rows)]
|
||||
|
||||
|
||||
def _numel(shape: list) -> int:
|
||||
n = 1
|
||||
for dim in shape:
|
||||
n *= dim
|
||||
return n
|
||||
|
||||
|
||||
def _linear(weight, bias, x):
|
||||
"""weight (out x in) @ x + bias -- radove poradie ako torch.nn.Linear."""
|
||||
return [sum(w * v for w, v in zip(row, x)) + b
|
||||
for row, b in zip(weight, bias)]
|
||||
|
||||
|
||||
def _relu(x):
|
||||
return [v if v > 0.0 else 0.0 for v in x]
|
||||
|
||||
|
||||
class PureNet:
|
||||
def __init__(self, payload: dict):
|
||||
if payload['obs_dim'] != OBS_DIM:
|
||||
raise ValueError(
|
||||
f'Vahy su pre obs_dim={payload["obs_dim"]}, kod ma {OBS_DIM} '
|
||||
'-- treba re-export z aktualneho checkpointu.'
|
||||
)
|
||||
w = payload['weights']
|
||||
self.trunk0_w = _unpack(w['trunk0_w'])
|
||||
self.trunk0_b = _unpack(w['trunk0_b'])
|
||||
self.trunk2_w = _unpack(w['trunk2_w'])
|
||||
self.trunk2_b = _unpack(w['trunk2_b'])
|
||||
self.guess_w = _unpack(w['guess_w'])
|
||||
self.guess_b = _unpack(w['guess_b'])
|
||||
self.play_w = _unpack(w['play_w'])
|
||||
self.play_b = _unpack(w['play_b'])
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: str = DEFAULT_WEIGHTS_PATH) -> 'PureNet':
|
||||
with open(path) as f:
|
||||
return cls(json.load(f))
|
||||
|
||||
def _trunk(self, obs):
|
||||
h = _relu(_linear(self.trunk0_w, self.trunk0_b, obs))
|
||||
return _relu(_linear(self.trunk2_w, self.trunk2_b, h))
|
||||
|
||||
def guess_logits(self, obs) -> list:
|
||||
return _linear(self.guess_w, self.guess_b, self._trunk(obs))
|
||||
|
||||
def play_logits(self, obs) -> list:
|
||||
return _linear(self.play_w, self.play_b, self._trunk(obs))
|
||||
|
||||
|
||||
def _masked_argmax(logits: list, mask: list) -> int:
|
||||
best, best_value = None, None
|
||||
for i, allowed in enumerate(mask):
|
||||
if allowed and (best is None or logits[i] > best_value):
|
||||
best, best_value = i, logits[i]
|
||||
return best
|
||||
|
||||
|
||||
class PureNeuralPlayer:
|
||||
"""Greedy hrac nad PureNet -- rovnake rozhranie a rovnake vstupy
|
||||
(observacia + maska) ako rl.policy_player.NeuralPlayer(greedy=True)."""
|
||||
|
||||
def __init__(self, net: PureNet):
|
||||
self.net = net
|
||||
|
||||
@classmethod
|
||||
def load(cls, path: str = DEFAULT_WEIGHTS_PATH) -> 'PureNeuralPlayer':
|
||||
return cls(PureNet.load(path))
|
||||
|
||||
def guess(self, rnd, seat: int) -> int:
|
||||
obs = encode_observation(rnd, seat)
|
||||
return _masked_argmax(self.net.guess_logits(obs), guess_mask(rnd))
|
||||
|
||||
def play(self, rnd, seat: int) -> int:
|
||||
obs = encode_observation(rnd, seat)
|
||||
return _masked_argmax(self.net.play_logits(obs), play_mask(rnd, seat))
|
||||
+131
@@ -0,0 +1,131 @@
|
||||
"""Self-play generator: jedna zdielana siet hra vsetkych 4 hracov v Round
|
||||
epizodach a zbiera trajektorie pre PPO (viz rl/DESIGN.md, sekcie 4-5).
|
||||
|
||||
Odmena je sparse a terminalna: kazde rozhodnutie hraca v kole (tip aj vsetky
|
||||
karty) dostane ako return jeho `points_summary` z konca kola, gamma = 1.
|
||||
|
||||
Masky sa ukladaju oddelene pre obe fazy (rozne velkosti akcneho priestoru);
|
||||
`phase_play` hovori, ktora hlava/maska pre dany krok plati.
|
||||
"""
|
||||
|
||||
from random import Random
|
||||
|
||||
import torch
|
||||
|
||||
from rl.encoding import N_GUESS_ACTIONS, N_PLAY_ACTIONS
|
||||
from rl.env import PHASE_PLAY, RoundEnv
|
||||
from rl.model import BridzikNet, mask_tensor, masked_categorical, obs_tensor
|
||||
from rl.players import HeuristicPlayer, RandomPlayer
|
||||
|
||||
# Returny sa skaluju do [0, 1] (max odmena je 10+8). Bez skalovania ma value
|
||||
# loss (MSE na 0-18) radovo vacsi gradient nez policy loss a cez zdielany
|
||||
# trup policy ucenie prevalcuje.
|
||||
REWARD_SCALE = 18.0
|
||||
|
||||
|
||||
def _assign_seats(rng: Random, mix_random: float, mix_heuristic: float,
|
||||
random_player, heuristic_player) -> dict:
|
||||
"""Obsadenie sedadiel pre jednu epizodu: None = siet, inak skriptovany
|
||||
supper. Aspon jedno sedadlo musi hrat siet (inak niet co zbierat)."""
|
||||
seats = {}
|
||||
for seat in range(4):
|
||||
roll = rng.random()
|
||||
if roll < mix_random:
|
||||
seats[seat] = random_player
|
||||
elif roll < mix_random + mix_heuristic:
|
||||
seats[seat] = heuristic_player
|
||||
else:
|
||||
seats[seat] = None
|
||||
if not any(p is None for p in seats.values()):
|
||||
seats[rng.randrange(4)] = None
|
||||
return seats
|
||||
|
||||
|
||||
def collect_episodes(net: BridzikNet, n_episodes: int, rng: Random,
|
||||
round_numbers: list = None, mix_random: float = 0.0,
|
||||
mix_heuristic: float = 0.0,
|
||||
heuristic_samples: int = 40) -> dict:
|
||||
"""Odohra `n_episodes` self-play kol a vrati batch tenzorov:
|
||||
|
||||
obs (N, OBS_DIM), phase_play (N,) bool, action (N,), logp (N,), value (N,),
|
||||
ret (N,), guess_mask (N, 9), play_mask (N, 32) -- maska nepatriacej fazy je
|
||||
pre dany krok cela False a pri update sa nepouzije.
|
||||
Navyse 'mean_points': priemerne body na sietove sedadlo a kolo.
|
||||
|
||||
Opponent mixing (robustnost na nie-self-play superov): s pravdepodobnostou
|
||||
`mix_random` / `mix_heuristic` hra sedadlo RandomPlayer / HeuristicPlayer
|
||||
namiesto siete. Tahy skriptovanych superov sa do batchu NEZAZNAMENAVAJU
|
||||
(nie su z trenovanej policy) -- superi len obsadzuju stol.
|
||||
"""
|
||||
env = RoundEnv(rng)
|
||||
random_player = RandomPlayer(rng)
|
||||
heuristic_player = HeuristicPlayer(rng, n_samples=heuristic_samples)
|
||||
mixing = mix_random > 0 or mix_heuristic > 0
|
||||
obs_l, phase_l, action_l, logp_l, value_l, ret_l = [], [], [], [], [], []
|
||||
gmask_l, pmask_l = [], []
|
||||
total_points = 0.0
|
||||
net_seat_rounds = 0
|
||||
|
||||
net.eval()
|
||||
with torch.no_grad():
|
||||
for _ in range(n_episodes):
|
||||
round_number = rng.choice(round_numbers) if round_numbers else None
|
||||
decision = env.reset(round_number)
|
||||
seats = _assign_seats(rng, mix_random, mix_heuristic,
|
||||
random_player, heuristic_player) if mixing \
|
||||
else {seat: None for seat in range(4)}
|
||||
net_seat_rounds += sum(1 for p in seats.values() if p is None)
|
||||
# indexy krokov sietovych sedadiel -- na priradenie returnu
|
||||
player_steps = {p: [] for p in range(4) if seats[p] is None}
|
||||
while True:
|
||||
opponent = seats[decision.player]
|
||||
if opponent is not None:
|
||||
# skriptovany supper: vykonaj tah, nic nezaznamenavaj
|
||||
if decision.phase == PHASE_PLAY:
|
||||
action_i = opponent.play(env.round, decision.player)
|
||||
else:
|
||||
action_i = opponent.guess(env.round, decision.player)
|
||||
else:
|
||||
obs = obs_tensor(decision.obs).unsqueeze(0)
|
||||
mask = mask_tensor(decision.mask).unsqueeze(0)
|
||||
guess_logits, play_logits, value = net(obs)
|
||||
is_play = decision.phase == PHASE_PLAY
|
||||
dist = masked_categorical(
|
||||
play_logits if is_play else guess_logits, mask
|
||||
)
|
||||
action = dist.sample()
|
||||
action_i = action.item()
|
||||
|
||||
player_steps[decision.player].append(len(obs_l))
|
||||
obs_l.append(decision.obs)
|
||||
phase_l.append(is_play)
|
||||
action_l.append(action_i)
|
||||
logp_l.append(dist.log_prob(action).item())
|
||||
value_l.append(value.item())
|
||||
ret_l.append(0.0) # doplni sa na konci kola
|
||||
if is_play:
|
||||
gmask_l.append([False] * N_GUESS_ACTIONS)
|
||||
pmask_l.append(decision.mask)
|
||||
else:
|
||||
gmask_l.append(decision.mask)
|
||||
pmask_l.append([False] * N_PLAY_ACTIONS)
|
||||
|
||||
decision, rewards, done = env.step(action_i)
|
||||
if done:
|
||||
for player, steps in player_steps.items():
|
||||
for i in steps:
|
||||
ret_l[i] = rewards[player] / REWARD_SCALE
|
||||
total_points += rewards[player]
|
||||
break
|
||||
|
||||
return {
|
||||
'obs': torch.tensor(obs_l, dtype=torch.float32),
|
||||
'phase_play': torch.tensor(phase_l, dtype=torch.bool),
|
||||
'action': torch.tensor(action_l, dtype=torch.long),
|
||||
'logp': torch.tensor(logp_l, dtype=torch.float32),
|
||||
'value': torch.tensor(value_l, dtype=torch.float32),
|
||||
'ret': torch.tensor(ret_l, dtype=torch.float32),
|
||||
'guess_mask': torch.tensor(gmask_l, dtype=torch.bool),
|
||||
'play_mask': torch.tensor(pmask_l, dtype=torch.bool),
|
||||
'mean_points': total_points / max(net_seat_rounds, 1),
|
||||
}
|
||||
+230
@@ -0,0 +1,230 @@
|
||||
"""Self-play PPO trening (viz rl/DESIGN.md, sekcie 4-6).
|
||||
|
||||
Slucka: nazbieraj self-play epizody -> PPO update -> kazdych par iteracii
|
||||
evaluacia GREEDY policy proti fixnym baseline-om (nahodny hrac, MC heuristika)
|
||||
z rl/players.py -- self-play reward sam o sebe nie je smerodajny (hra nie je
|
||||
zero-sum, protihrac sa hybe spolu so sietou).
|
||||
|
||||
Spustenie:
|
||||
py -m rl.train --iterations 200 --episodes 512
|
||||
py -m rl.train --resume rl/checkpoints/latest.pt # pokracovanie
|
||||
|
||||
Checkpointy: rl/checkpoints/latest.pt (kazdu iteraciu) + best.pt (najlepsi
|
||||
priemer bodov proti heuristikam). Metriky sa pripisuju do rl/runs/train_log.csv.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import os
|
||||
import time
|
||||
from random import Random
|
||||
|
||||
import torch
|
||||
|
||||
from rl.evaluate import evaluate
|
||||
from rl.model import BridzikNet, masked_categorical
|
||||
from rl.players import HeuristicPlayer, RandomPlayer
|
||||
from rl.policy_player import NeuralPlayer
|
||||
from rl.selfplay import collect_episodes
|
||||
|
||||
CHECKPOINT_DIR = os.path.join('rl', 'checkpoints')
|
||||
RUNS_DIR = os.path.join('rl', 'runs')
|
||||
|
||||
|
||||
def ppo_update(net: BridzikNet, optimizer: torch.optim.Optimizer, batch: dict,
|
||||
clip: float = 0.2, epochs: int = 4, minibatch: int = 1024,
|
||||
vf_coef: float = 1.0, ent_coef: float = 0.01,
|
||||
max_grad_norm: float = 0.5) -> dict:
|
||||
"""Standardny clipped-PPO krok nad batchom zo self-play.
|
||||
|
||||
Advantage sa standardizuje per batch (bod 7 v DESIGN.md -- odmeny 10-18 sa
|
||||
lisia medzi kolami a zvysovali by varianciu gradientu). Guess a play kroky
|
||||
zdielaju trup aj value hlavu, policy loss ide vzdy cez hlavu svojej fazy.
|
||||
"""
|
||||
n = batch['obs'].shape[0]
|
||||
adv = batch['ret'] - batch['value']
|
||||
adv = (adv - adv.mean()) / (adv.std() + 1e-8)
|
||||
|
||||
net.train()
|
||||
stats = {'policy_loss': 0.0, 'value_loss': 0.0, 'entropy': 0.0, 'updates': 0}
|
||||
for _ in range(epochs):
|
||||
perm = torch.randperm(n)
|
||||
for start in range(0, n, minibatch):
|
||||
idx = perm[start:start + minibatch]
|
||||
obs = batch['obs'][idx]
|
||||
guess_logits, play_logits, value = net(obs)
|
||||
|
||||
is_play = batch['phase_play'][idx]
|
||||
logp_new = torch.empty_like(batch['logp'][idx])
|
||||
entropy = torch.empty_like(logp_new)
|
||||
for phase_sel, logits, mask_key in (
|
||||
(~is_play, guess_logits, 'guess_mask'),
|
||||
(is_play, play_logits, 'play_mask'),
|
||||
):
|
||||
if not bool(phase_sel.any()):
|
||||
continue
|
||||
dist = masked_categorical(
|
||||
logits[phase_sel], batch[mask_key][idx][phase_sel]
|
||||
)
|
||||
logp_new[phase_sel] = dist.log_prob(batch['action'][idx][phase_sel])
|
||||
entropy[phase_sel] = dist.entropy()
|
||||
|
||||
ratio = torch.exp(logp_new - batch['logp'][idx])
|
||||
mb_adv = adv[idx]
|
||||
policy_loss = -torch.min(
|
||||
ratio * mb_adv,
|
||||
torch.clamp(ratio, 1 - clip, 1 + clip) * mb_adv,
|
||||
).mean()
|
||||
value_loss = (value - batch['ret'][idx]).pow(2).mean()
|
||||
loss = policy_loss + vf_coef * value_loss - ent_coef * entropy.mean()
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
torch.nn.utils.clip_grad_norm_(net.parameters(), max_grad_norm)
|
||||
optimizer.step()
|
||||
|
||||
stats['policy_loss'] += policy_loss.item()
|
||||
stats['value_loss'] += value_loss.item()
|
||||
stats['entropy'] += entropy.mean().item()
|
||||
stats['updates'] += 1
|
||||
|
||||
for key in ('policy_loss', 'value_loss', 'entropy'):
|
||||
stats[key] /= max(stats['updates'], 1)
|
||||
return stats
|
||||
|
||||
|
||||
def evaluate_against_baselines(net: BridzikNet, n_rounds: int, rng: Random,
|
||||
mc_samples: int = 60) -> dict:
|
||||
"""Greedy siet na sedadle 0 vs 3x random a vs 3x heuristika."""
|
||||
neural = NeuralPlayer(net, greedy=True)
|
||||
vs_random = evaluate(
|
||||
[neural] + [RandomPlayer(rng) for _ in range(3)], n_rounds, rng
|
||||
)[0]
|
||||
vs_heuristic = evaluate(
|
||||
[neural] + [HeuristicPlayer(rng, n_samples=mc_samples) for _ in range(3)],
|
||||
n_rounds, rng,
|
||||
)[0]
|
||||
return {
|
||||
'vs_random_points': vs_random['avg_points'],
|
||||
'vs_random_hit': vs_random['hit_rate'],
|
||||
'vs_heuristic_points': vs_heuristic['avg_points'],
|
||||
'vs_heuristic_hit': vs_heuristic['hit_rate'],
|
||||
}
|
||||
|
||||
|
||||
def save_checkpoint(net: BridzikNet, hidden: int, path: str) -> None:
|
||||
torch.save({'hidden': hidden, 'state_dict': net.state_dict()}, path)
|
||||
|
||||
|
||||
def load_checkpoint(path: str) -> BridzikNet:
|
||||
"""Nacita checkpoint; podporuje aj stary format (bare state_dict)."""
|
||||
payload = torch.load(path, map_location='cpu')
|
||||
if isinstance(payload, dict) and 'state_dict' in payload:
|
||||
net = BridzikNet(hidden=payload['hidden'])
|
||||
net.load_state_dict(payload['state_dict'])
|
||||
else:
|
||||
net = BridzikNet()
|
||||
net.load_state_dict(payload)
|
||||
return net
|
||||
|
||||
|
||||
def train(iterations: int, episodes: int, lr: float, seed: int,
|
||||
eval_every: int, eval_rounds: int, resume: str = None,
|
||||
hidden: int = 384, ent_coef_start: float = 0.01,
|
||||
ent_coef_final: float = 0.001, lr_final_frac: float = 0.1,
|
||||
mix_random: float = 0.0, mix_heuristic: float = 0.0):
|
||||
os.makedirs(CHECKPOINT_DIR, exist_ok=True)
|
||||
os.makedirs(RUNS_DIR, exist_ok=True)
|
||||
log_path = os.path.join(RUNS_DIR, 'train_log.csv')
|
||||
log_exists = os.path.exists(log_path)
|
||||
|
||||
torch.manual_seed(seed)
|
||||
rng = Random(seed)
|
||||
if resume:
|
||||
net = load_checkpoint(resume)
|
||||
hidden = net.trunk[0].out_features
|
||||
print(f'Pokracujem z checkpointu {resume} (hidden={hidden})')
|
||||
else:
|
||||
net = BridzikNet(hidden=hidden)
|
||||
optimizer = torch.optim.Adam(net.parameters(), lr=lr)
|
||||
|
||||
best_vs_heuristic = float('-inf')
|
||||
with open(log_path, 'a', newline='') as log_file:
|
||||
log = csv.writer(log_file)
|
||||
if not log_exists:
|
||||
log.writerow(['iteration', 'selfplay_points', 'policy_loss',
|
||||
'value_loss', 'entropy', 'vs_random_points',
|
||||
'vs_random_hit', 'vs_heuristic_points',
|
||||
'vs_heuristic_hit', 'seconds'])
|
||||
|
||||
for iteration in range(1, iterations + 1):
|
||||
started = time.time()
|
||||
# linearny decay: lr klesa k lr*lr_final_frac, entropny bonus
|
||||
# k ent_coef_final -- policy sa ku koncu behu moze doostrit
|
||||
frac = 1 - (iteration - 1) / max(iterations - 1, 1)
|
||||
for group in optimizer.param_groups:
|
||||
group['lr'] = lr * (lr_final_frac + (1 - lr_final_frac) * frac)
|
||||
ent_coef = ent_coef_final + (ent_coef_start - ent_coef_final) * frac
|
||||
|
||||
batch = collect_episodes(net, episodes, rng,
|
||||
mix_random=mix_random,
|
||||
mix_heuristic=mix_heuristic)
|
||||
stats = ppo_update(net, optimizer, batch, ent_coef=ent_coef)
|
||||
save_checkpoint(net, hidden, os.path.join(CHECKPOINT_DIR, 'latest.pt'))
|
||||
|
||||
row = [iteration, f'{batch["mean_points"]:.3f}',
|
||||
f'{stats["policy_loss"]:.4f}', f'{stats["value_loss"]:.2f}',
|
||||
f'{stats["entropy"]:.3f}']
|
||||
line = (f'it {iteration:4d} | self-play {batch["mean_points"]:5.2f} '
|
||||
f'b/kolo | pi {stats["policy_loss"]:+.4f} '
|
||||
f'| V {stats["value_loss"]:7.2f} | H {stats["entropy"]:.3f}')
|
||||
|
||||
if iteration % eval_every == 0 or iteration == iterations:
|
||||
ev = evaluate_against_baselines(net, eval_rounds, rng)
|
||||
row += [f'{ev["vs_random_points"]:.3f}', f'{ev["vs_random_hit"]:.3f}',
|
||||
f'{ev["vs_heuristic_points"]:.3f}', f'{ev["vs_heuristic_hit"]:.3f}']
|
||||
line += (f' | vs random {ev["vs_random_points"]:5.2f} '
|
||||
f'({100 * ev["vs_random_hit"]:.0f} %)'
|
||||
f' | vs heur {ev["vs_heuristic_points"]:5.2f} '
|
||||
f'({100 * ev["vs_heuristic_hit"]:.0f} %)')
|
||||
if ev['vs_heuristic_points'] > best_vs_heuristic:
|
||||
best_vs_heuristic = ev['vs_heuristic_points']
|
||||
save_checkpoint(net, hidden,
|
||||
os.path.join(CHECKPOINT_DIR, 'best.pt'))
|
||||
line += ' *best*'
|
||||
else:
|
||||
row += ['', '', '', '']
|
||||
|
||||
row.append(f'{time.time() - started:.1f}')
|
||||
log.writerow(row)
|
||||
log_file.flush()
|
||||
print(line)
|
||||
|
||||
return net
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description='Self-play PPO trening bridzik siete')
|
||||
parser.add_argument('--iterations', type=int, default=200)
|
||||
parser.add_argument('--episodes', type=int, default=512,
|
||||
help='self-play kol na iteraciu')
|
||||
parser.add_argument('--lr', type=float, default=3e-4)
|
||||
parser.add_argument('--seed', type=int, default=1)
|
||||
parser.add_argument('--eval-every', type=int, default=10)
|
||||
parser.add_argument('--eval-rounds', type=int, default=400)
|
||||
parser.add_argument('--hidden', type=int, default=384,
|
||||
help='sirka skrytych vrstiev trupu')
|
||||
parser.add_argument('--resume', type=str, default=None,
|
||||
help='cesta k checkpointu (.pt) na pokracovanie')
|
||||
parser.add_argument('--mix-random', type=float, default=0.0,
|
||||
help='pravdepodobnost RandomPlayer sedadla v epizode')
|
||||
parser.add_argument('--mix-heuristic', type=float, default=0.0,
|
||||
help='pravdepodobnost HeuristicPlayer sedadla v epizode')
|
||||
args = parser.parse_args()
|
||||
train(args.iterations, args.episodes, args.lr, args.seed,
|
||||
args.eval_every, args.eval_rounds, args.resume, hidden=args.hidden,
|
||||
mix_random=args.mix_random, mix_heuristic=args.mix_heuristic)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,218 @@
|
||||
"""Testy in-process botov (api/bots.py + tahova slucka v api/__init__.py).
|
||||
|
||||
Rovnaky setup ako tests/test_history.py: docasny SQLite subor, env pred
|
||||
importom. Socket.IO emity idu do prazdnych roomov (ziadny klient), takze
|
||||
handlery a slucka sa daju volat priamo bez klienta.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
import uuid
|
||||
|
||||
# Nastav DB/ENCRYPTION_KEY PRED importom db/api modulov.
|
||||
_DB_FILE = os.path.join(tempfile.gettempdir(), f"bridzik_test_{uuid.uuid4().hex}.db")
|
||||
os.environ["DATABASE_URL"] = "sqlite+aiosqlite:///" + _DB_FILE.replace("\\", "/")
|
||||
|
||||
from cryptography.fernet import Fernet # noqa: E402
|
||||
|
||||
os.environ.setdefault("ENCRYPTION_KEY", Fernet.generate_key().decode())
|
||||
|
||||
from random import Random # noqa: E402
|
||||
|
||||
import api # noqa: E402
|
||||
from api import auth, bots, history # noqa: E402
|
||||
from db.db import init_db # noqa: E402
|
||||
from rl.players import HeuristicPlayer, RandomPlayer # noqa: E402
|
||||
|
||||
|
||||
def run(coro):
|
||||
return asyncio.run(coro)
|
||||
|
||||
|
||||
def _make_bot_accounts(n):
|
||||
exclude = set()
|
||||
accounts = []
|
||||
for _ in range(n):
|
||||
acc = run(bots.ensure_bot_account("random", exclude))
|
||||
exclude.add(acc["player_id"])
|
||||
accounts.append(acc)
|
||||
return accounts
|
||||
|
||||
|
||||
def _make_game(seat_accounts, brains):
|
||||
"""Postavi zacatu in-memory hru + Game riadok v DB."""
|
||||
gid = str(uuid.uuid4())
|
||||
game = api.Game(gid, "test")
|
||||
for seat, acc in enumerate(seat_accounts):
|
||||
player = api.Player(None, acc["username"], seat, acc["player_id"])
|
||||
if brains[seat] is not None:
|
||||
player.is_bot = True
|
||||
player.brain = brains[seat]
|
||||
player.connected = True
|
||||
else:
|
||||
player.connected = False
|
||||
game.players.append(player)
|
||||
game.start()
|
||||
api.games[gid] = game
|
||||
run(history.record_game_started(
|
||||
gid, "test", [acc["player_id"] for acc in seat_accounts]
|
||||
))
|
||||
return game
|
||||
|
||||
|
||||
class BotAccountCase(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
run(init_db())
|
||||
api.BOT_MOVE_DELAY_SECONDS = 0
|
||||
api.TRICK_SWEEP_SECONDS = 0
|
||||
|
||||
def test_username_conventions(self):
|
||||
self.assertTrue(bots.is_bot_username("bot:heuristic-1"))
|
||||
self.assertFalse(bots.is_bot_username("alice"))
|
||||
self.assertEqual(bots.kind_of("bot:heuristic-2"), "heuristic")
|
||||
self.assertEqual(bots.kind_of("bot:random-1"), "random")
|
||||
self.assertEqual(bots.kind_of("bot:neural-1"), "neural")
|
||||
self.assertEqual(bots.kind_of("bot:nezmysel-9"), bots.DEFAULT_KIND)
|
||||
self.assertIsInstance(bots.make_brain("bot:heuristic-1"), HeuristicPlayer)
|
||||
self.assertIsInstance(bots.make_brain("bot:random-3"), RandomPlayer)
|
||||
|
||||
@unittest.skipUnless(bots.neural_available(),
|
||||
'chyba export vah (py -m rl.export)')
|
||||
def test_neural_kind(self):
|
||||
from rl.pure_net import PureNeuralPlayer
|
||||
self.assertIn("neural", bots.available_kinds())
|
||||
brain = bots.make_brain("bot:neural-1")
|
||||
self.assertIsInstance(brain, PureNeuralPlayer)
|
||||
# zdielana instancia PureNet (vahy sa nacitavaju len raz)
|
||||
self.assertIs(brain.net, bots.make_brain("bot:neural-2").net)
|
||||
|
||||
def test_ensure_bot_account_reuse_and_exclude(self):
|
||||
first = run(bots.ensure_bot_account("heuristic", set()))
|
||||
self.assertTrue(first["username"].startswith("bot:heuristic-"))
|
||||
# bez vylucenia sa ucet recykluje
|
||||
again = run(bots.ensure_bot_account("heuristic", set()))
|
||||
self.assertEqual(first["player_id"], again["player_id"])
|
||||
# s vylucenim vznikne dalsi ucet s inym ID
|
||||
second = run(bots.ensure_bot_account("heuristic", {first["player_id"]}))
|
||||
self.assertNotEqual(first["player_id"], second["player_id"])
|
||||
self.assertNotEqual(first["username"], second["username"])
|
||||
|
||||
def test_bot_account_cannot_be_hijacked(self):
|
||||
acc = run(bots.ensure_bot_account("heuristic", set()))
|
||||
# registracia mena zlyha -- ucet sa netvari ako nedokoncena registracia
|
||||
with self.assertRaises(auth.AuthError):
|
||||
run(auth.register_account(acc["username"]))
|
||||
# login zlyha na kode (secret nikto nepozna), NIE na RegistrationIncomplete
|
||||
with self.assertRaises(auth.AuthError) as ctx:
|
||||
run(auth.login(acc["username"], "000000"))
|
||||
self.assertNotIsInstance(ctx.exception, auth.RegistrationIncomplete)
|
||||
|
||||
|
||||
class BotTurnLoopCase(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
run(init_db())
|
||||
api.BOT_MOVE_DELAY_SECONDS = 0
|
||||
api.TRICK_SWEEP_SECONDS = 0
|
||||
|
||||
def setUp(self):
|
||||
api.games.clear()
|
||||
api.sessions.clear()
|
||||
api.accounts.clear()
|
||||
|
||||
def test_four_bots_play_whole_game(self):
|
||||
accounts = _make_bot_accounts(4)
|
||||
brains = [RandomPlayer(Random(seat)) for seat in range(4)]
|
||||
game = _make_game(accounts, brains)
|
||||
|
||||
run(api._run_bot_turns(game.gid))
|
||||
|
||||
self.assertTrue(game.bridzik_core.is_completed())
|
||||
# cela hra je zapisana: 4 serie x 8 kol, ended_at nastaveny
|
||||
detail = run(history.get_game_detail(game.gid))
|
||||
self.assertIsNotNone(detail["ended_at"])
|
||||
self.assertEqual(len(detail["rounds"]), history.FULL_GAME_ROUNDS * 4)
|
||||
|
||||
def test_bots_stop_at_human_turn(self):
|
||||
accounts = _make_bot_accounts(3)
|
||||
# sedadlo 0 = clovek; identitu v DB mu robi dalsi (nepouzity) boti
|
||||
# ucet -- pre historiu je to len player_id, wrapper bez mozgu = clovek
|
||||
human_acc = run(bots.ensure_bot_account(
|
||||
"random", {a["player_id"] for a in accounts}
|
||||
))
|
||||
game = _make_game(
|
||||
[human_acc] + accounts,
|
||||
[None, HeuristicPlayer(Random(1), n_samples=20),
|
||||
RandomPlayer(Random(2)), RandomPlayer(Random(3))],
|
||||
)
|
||||
core = game.bridzik_core
|
||||
rnd = core.series[-1].get_last_round()
|
||||
|
||||
async def scenario():
|
||||
# na tahu je clovek (first_player serie 0 je sedadlo 0) -> boti nic
|
||||
await api._run_bot_turns(game.gid)
|
||||
self.assertEqual(len(rnd.guesses), 0)
|
||||
# clovek tipne -> boti dotipuju a hraju az po dalsi tah cloveka
|
||||
core.add_player_guess(0, 1)
|
||||
await api._run_bot_turns(game.gid)
|
||||
|
||||
run(scenario())
|
||||
self.assertTrue(rnd.is_guessing_completed())
|
||||
self.assertEqual(rnd.get_active_player(), 0)
|
||||
|
||||
def test_add_and_remove_bot_handlers(self):
|
||||
async def scenario():
|
||||
gid = str(uuid.uuid4())
|
||||
api.games[gid] = api.Game(gid, "lobby-test")
|
||||
host = api.Player("sid-host", "hostiteľ", 0, 999_100)
|
||||
api.games[gid].players.append(host)
|
||||
api.sessions["sid-host"] = {"gid": gid, "order": 0}
|
||||
api.sessions["sid-guest"] = {"gid": gid, "order": 1}
|
||||
|
||||
# nehostitel nesmie pridat bota
|
||||
await api.add_bot("sid-guest", gid)
|
||||
self.assertEqual(len(api.games[gid].players), 1)
|
||||
|
||||
# hostitel prida dvoch botov -> rozne ucty, najnizsie volne sedadla
|
||||
await api.add_bot("sid-host", gid)
|
||||
await api.add_bot("sid-host", gid, "random")
|
||||
players = api.games[gid].players
|
||||
self.assertEqual(len(players), 3)
|
||||
bots_added = [p for p in players if p.is_bot]
|
||||
self.assertEqual(len(bots_added), 2)
|
||||
self.assertEqual({p.order for p in bots_added}, {1, 2})
|
||||
self.assertNotEqual(bots_added[0].player_id, bots_added[1].player_id)
|
||||
self.assertIsNotNone(bots_added[0].brain)
|
||||
|
||||
# remove_bot: odmietne cloveka, odoberie bota
|
||||
await api.remove_bot("sid-host", gid, 0)
|
||||
self.assertEqual(len(api.games[gid].players), 3)
|
||||
await api.remove_bot("sid-host", gid, 1)
|
||||
self.assertEqual(len(api.games[gid].players), 2)
|
||||
self.assertIsNone(api.games[gid].player_by_order(1))
|
||||
|
||||
run(scenario())
|
||||
|
||||
def test_restore_marks_bots(self):
|
||||
from bridzik import Bridzik
|
||||
info = {
|
||||
"gid": str(uuid.uuid4()),
|
||||
"name": "obnova",
|
||||
"seats": [(1, "alice"), (2, "bot:heuristic-1"),
|
||||
(3, "bot:random-1"), (4, "bob")],
|
||||
"core": Bridzik(),
|
||||
}
|
||||
game = api._load_game_into_memory(info)
|
||||
self.assertFalse(game.players[0].is_bot)
|
||||
self.assertFalse(game.players[0].connected)
|
||||
self.assertTrue(game.players[1].is_bot)
|
||||
self.assertTrue(game.players[1].connected)
|
||||
self.assertIsInstance(game.players[1].brain, HeuristicPlayer)
|
||||
self.assertIsInstance(game.players[2].brain, RandomPlayer)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main(verbosity=2)
|
||||
@@ -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)
|
||||
@@ -121,6 +121,16 @@ class StashCase(unittest.TestCase):
|
||||
self.assertRaises(BridzikException, s.get_active_player)
|
||||
|
||||
|
||||
class CardCase(unittest.TestCase):
|
||||
def test_hashable_consistent_with_eq(self):
|
||||
# vlastne __eq__ nesmie zrusit hashovatelnost (dict/set pouzitie)
|
||||
heart_7 = Card(Card_colors['HEARTS'], Card_values['C7'])
|
||||
self.assertEqual(hash(heart_7), hash(Card(Card_colors['HEARTS'], Card_values['C7'])))
|
||||
self.assertEqual(len(set(cards)), 32)
|
||||
self.assertEqual({Card_colors['HEARTS']: 1}[Card_colors['HEARTS']], 1)
|
||||
self.assertEqual({Card_values['ACE']: 1}[Card_values['ACE']], 1)
|
||||
|
||||
|
||||
class RoundCase(unittest.TestCase):
|
||||
def test_round_constructor(self):
|
||||
self.assertRaises(BridzikException, Round, round_number=8, first_player=0)
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
"""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)
|
||||
@@ -0,0 +1,316 @@
|
||||
import unittest
|
||||
from random import Random
|
||||
|
||||
from bridzik import cards, Card, Card_colors, Card_values, Round
|
||||
from rl.encoding import card_index, index_card, legal_cards
|
||||
from rl.env import Decision, PHASE_GUESS, PHASE_PLAY, RoundEnv
|
||||
from rl.evaluate import evaluate, play_round
|
||||
from rl.players import (
|
||||
HeuristicPlayer, McPlayer, RandomPlayer,
|
||||
_beats, _current_best, deal_consistent, deduce_voids,
|
||||
mc_guess_distribution, simulate_tricks,
|
||||
)
|
||||
|
||||
|
||||
class RoundEnvCase(unittest.TestCase):
|
||||
def test_episode_structure(self):
|
||||
env = RoundEnv(Random(42))
|
||||
decision = env.reset(round_number=6, first_player=1)
|
||||
rng = Random(0)
|
||||
|
||||
# prve 4 rozhodnutia su tipy, v poradi od first_player
|
||||
expected_guessers = [1, 2, 3, 0]
|
||||
for expected in expected_guessers:
|
||||
self.assertIsInstance(decision, Decision)
|
||||
self.assertEqual(decision.phase, PHASE_GUESS)
|
||||
self.assertEqual(decision.player, expected)
|
||||
action = rng.choice([g for g in range(9) if decision.mask[g]])
|
||||
decision, rewards, done = env.step(action)
|
||||
self.assertIsNone(rewards)
|
||||
self.assertFalse(done)
|
||||
|
||||
# potom hracie rozhodnutia az po terminal: 2 karty x 4 hraci
|
||||
steps = 0
|
||||
while True:
|
||||
self.assertEqual(decision.phase, PHASE_PLAY)
|
||||
self.assertEqual(decision.player, env.round.get_active_player())
|
||||
action = rng.choice([i for i in range(32) if decision.mask[i]])
|
||||
decision, rewards, done = env.step(action)
|
||||
steps += 1
|
||||
if done:
|
||||
break
|
||||
self.assertEqual(steps, 8)
|
||||
self.assertIsNone(decision)
|
||||
self.assertEqual(rewards, env.round.get_points_summary())
|
||||
self.assertEqual(len(rewards), 4)
|
||||
|
||||
# po done sa step neda volat, reset zacne novu epizodu
|
||||
self.assertRaises(RuntimeError, env.step, 0)
|
||||
self.assertIsInstance(env.reset(), Decision)
|
||||
|
||||
def test_reset_samples_round_and_seat(self):
|
||||
env = RoundEnv(Random(7))
|
||||
seen_rounds, seen_seats = set(), set()
|
||||
for _ in range(100):
|
||||
env.reset()
|
||||
seen_rounds.add(env.round.round_number)
|
||||
seen_seats.add(env.round.first_player)
|
||||
self.assertEqual(seen_rounds, set(range(8)))
|
||||
self.assertEqual(seen_seats, set(range(4)))
|
||||
|
||||
def test_deterministic_with_seed(self):
|
||||
rewards = []
|
||||
for _ in range(2):
|
||||
env = RoundEnv(Random(123))
|
||||
players = [RandomPlayer(Random(5)) for _ in range(4)]
|
||||
rewards.append(play_round(players, env, round_number=0))
|
||||
self.assertEqual(rewards[0], rewards[1])
|
||||
|
||||
|
||||
class SimulationHelpersCase(unittest.TestCase):
|
||||
def test_beats(self):
|
||||
heart_7 = Card(Card_colors['HEARTS'], Card_values['C7'])
|
||||
heart_8 = Card(Card_colors['HEARTS'], Card_values['C8'])
|
||||
leaves_ace = Card(Card_colors['LEAVES'], Card_values['ACE'])
|
||||
leaves_king = Card(Card_colors['LEAVES'], Card_values['KING'])
|
||||
bells_ace = Card(Card_colors['BELLS'], Card_values['ACE'])
|
||||
|
||||
self.assertTrue(_beats(leaves_ace, leaves_king)) # vyssia vo farbe
|
||||
self.assertFalse(_beats(leaves_king, leaves_ace))
|
||||
self.assertTrue(_beats(heart_7, leaves_ace)) # tromf bije farbu
|
||||
self.assertFalse(_beats(bells_ace, leaves_king)) # cudzia farba neberie
|
||||
self.assertTrue(_beats(heart_8, heart_7)) # tromfy medzi sebou
|
||||
self.assertFalse(_beats(leaves_ace, heart_7)) # farba nebije tromf
|
||||
|
||||
def test_current_best_tracks_stash(self):
|
||||
from bridzik import Stash
|
||||
leaves_7 = Card(Card_colors['LEAVES'], Card_values['C7'])
|
||||
leaves_ace = Card(Card_colors['LEAVES'], Card_values['ACE'])
|
||||
heart_7 = Card(Card_colors['HEARTS'], Card_values['C7'])
|
||||
|
||||
self.assertIsNone(_current_best(None))
|
||||
s = Stash(0)
|
||||
self.assertIsNone(_current_best(s))
|
||||
s.add_card(0, leaves_7)
|
||||
self.assertEqual(_current_best(s), leaves_7)
|
||||
s.add_card(1, leaves_ace)
|
||||
self.assertEqual(_current_best(s), leaves_ace)
|
||||
s.add_card(2, heart_7)
|
||||
self.assertEqual(_current_best(s), heart_7)
|
||||
|
||||
def test_simulate_tricks_consumes_hands(self):
|
||||
rng = Random(3)
|
||||
deck = list(cards)
|
||||
rng.shuffle(deck)
|
||||
hands = {seat: deck[seat * 8:(seat + 1) * 8] for seat in range(4)}
|
||||
tricks = simulate_tricks(hands, leader=2, rng=rng)
|
||||
self.assertEqual(sum(tricks), 8)
|
||||
for seat in range(4):
|
||||
self.assertEqual(hands[seat], [])
|
||||
|
||||
|
||||
class HeuristicPlayerCase(unittest.TestCase):
|
||||
@staticmethod
|
||||
def _round_with_hand(player0_hand):
|
||||
# deterministicke rozdanie: player0_hand ide hracovi 0, zvysok dalej;
|
||||
# deal_starting_cards najprv zahodi 4*round_number kariet, preto
|
||||
# treba ruku umiestnit az ZA odkladaciu kopu
|
||||
round_number = 8 - len(player0_hand)
|
||||
rest = [c for c in cards if c not in player0_hand]
|
||||
skip = 4 * round_number
|
||||
deck = rest[:skip] + list(player0_hand) + rest[skip:]
|
||||
return Round(round_number, 0, deck, shuffler=lambda l: None)
|
||||
|
||||
def test_mc_guess_all_hearts_is_certain(self):
|
||||
# 8 cerveni = tromfy beru kazdu kopku bez ohladu na rozdanie a hru
|
||||
all_hearts = [Card(Card_colors['HEARTS'], v) for v in Card_values]
|
||||
r = self._round_with_hand(all_hearts)
|
||||
counts = mc_guess_distribution(r, 0, n_samples=30, rng=Random(1))
|
||||
self.assertEqual(counts, {8: 30})
|
||||
self.assertEqual(HeuristicPlayer(Random(1), n_samples=30).guess(r, 0), 8)
|
||||
|
||||
def test_mc_guess_weak_hand_low(self):
|
||||
# dve najnizsie necervene karty -> tip 0 s prehladom
|
||||
weak = [Card(Card_colors['LEAVES'], Card_values['C7']),
|
||||
Card(Card_colors['BELLS'], Card_values['C7'])]
|
||||
r = self._round_with_hand(weak)
|
||||
self.assertEqual(HeuristicPlayer(Random(2), n_samples=60).guess(r, 0), 0)
|
||||
|
||||
def test_mc_guess_respects_mask(self):
|
||||
# posledny tipujuci: zakazana hodnota nesmie byt vratena, ani ked
|
||||
# je modom rozdelenia
|
||||
all_hearts = [Card(Card_colors['HEARTS'], v) for v in Card_values]
|
||||
r = self._round_with_hand(all_hearts)
|
||||
r.add_player_guess(0, 0)
|
||||
r.add_player_guess(1, 0)
|
||||
r.add_player_guess(2, 0)
|
||||
# zakazany tip pre hraca 3 je 8; jeho ruka je nahodna, ale nech by
|
||||
# simulacia vratila cokolvek, vysledok musi byt legalny
|
||||
guess = HeuristicPlayer(Random(3), n_samples=20).guess(r, 3)
|
||||
self.assertNotEqual(guess, 8)
|
||||
self.assertIn(guess, range(8))
|
||||
|
||||
def test_play_takes_trick_when_needed(self):
|
||||
hand = [Card(Card_colors['LEAVES'], Card_values['ACE']),
|
||||
Card(Card_colors['LEAVES'], Card_values['C7']),
|
||||
Card(Card_colors['BELLS'], Card_values['C7'])]
|
||||
r = self._round_with_hand(hand)
|
||||
r.add_player_guess(0, 3) # najvyssi tip -> hrac 0 vynasa
|
||||
r.add_player_guess(1, 0)
|
||||
r.add_player_guess(2, 0)
|
||||
r.add_player_guess(3, 1)
|
||||
# hrac 0 potrebuje kopky -> vynasa najsilnejsiu kartu (LEAVES ACE)
|
||||
action = HeuristicPlayer(Random(4)).play(r, 0)
|
||||
self.assertEqual(index_card(action), hand[0])
|
||||
|
||||
def test_play_ducks_when_satisfied(self):
|
||||
hand = [Card(Card_colors['LEAVES'], Card_values['ACE']),
|
||||
Card(Card_colors['LEAVES'], Card_values['C7']),
|
||||
Card(Card_colors['BELLS'], Card_values['C7'])]
|
||||
r = self._round_with_hand(hand)
|
||||
r.add_player_guess(0, 0) # hrac 0 nechce ziadnu kopku
|
||||
r.add_player_guess(1, 2) # najvyssi tip -> vynasa hrac 1
|
||||
r.add_player_guess(2, 0)
|
||||
r.add_player_guess(3, 0)
|
||||
first_card = legal_cards(r.player_cards[1], None)[0]
|
||||
r.play_card(1, first_card)
|
||||
action = HeuristicPlayer(Random(5)).play(r, 2)
|
||||
# legalnost staci overit enginom; strategiu netestujeme natvrdo,
|
||||
# lebo zavisi od nahodnej ruky hraca 2
|
||||
r.play_card(2, index_card(action))
|
||||
|
||||
def test_play_duck_scenario_deterministic(self):
|
||||
# hrac 0 tipol 0, ma na ruke LEAVES ACE aj C7; kopku vedie LEAVES C8
|
||||
# -> musi priznat farbu a spravne je podliezt (C7), nie zobrat esom
|
||||
hand0 = [Card(Card_colors['LEAVES'], Card_values['ACE']),
|
||||
Card(Card_colors['LEAVES'], Card_values['C7'])]
|
||||
hand1 = [Card(Card_colors['LEAVES'], Card_values['C8']),
|
||||
Card(Card_colors['LEAVES'], Card_values['C9'])]
|
||||
rest = [c for c in cards if c not in hand0 + hand1]
|
||||
deck = rest[:24] + hand0 + hand1 + rest[24:] # 24 = odkladacia kopa
|
||||
r = Round(6, 0, deck, shuffler=lambda l: None)
|
||||
r.add_player_guess(0, 0)
|
||||
r.add_player_guess(1, 2) # vynasa hrac 1
|
||||
r.add_player_guess(2, 0)
|
||||
r.add_player_guess(3, 1) # 0+2+0+0 by bol zakazany sucet (2 kopky)
|
||||
r.play_card(1, hand1[0])
|
||||
action = HeuristicPlayer(Random(6)).play(r, 0)
|
||||
self.assertEqual(index_card(action), hand0[1])
|
||||
|
||||
|
||||
class VoidDeductionCase(unittest.TestCase):
|
||||
def test_deduce_voids_from_stash(self):
|
||||
# hrac 0 vynasa zelen; 1 prizna farbu (nic), 2 tromfne cervenou
|
||||
# (void zelen), 3 hodi gulu (void zelen AJ cerven)
|
||||
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) # najvyssi tip -> vynasa 0
|
||||
r.add_player_guess(1, 0)
|
||||
r.add_player_guess(2, 0)
|
||||
r.add_player_guess(3, 1)
|
||||
|
||||
self.assertEqual(deduce_voids(r), {0: set(), 1: set(), 2: set(), 3: set()})
|
||||
r.play_card(0, hand0[0])
|
||||
r.play_card(1, hand1[0]) # priznal farbu -> nic
|
||||
r.play_card(2, hand2[0]) # cerven -> void zelen
|
||||
r.play_card(3, hand3[0]) # gula -> void zelen aj cerven
|
||||
voids = deduce_voids(r)
|
||||
self.assertEqual(voids[0], set()) # vynasajuci neprezradza nic
|
||||
self.assertEqual(voids[1], set())
|
||||
self.assertEqual(voids[2], {Card_colors['LEAVES']})
|
||||
self.assertEqual(voids[3], {Card_colors['LEAVES'], Card_colors['HEARTS']})
|
||||
|
||||
def test_deduced_voids_never_contradict_hands(self):
|
||||
# fuzz: dedukovany void NIKDY neprotireci realnej ruke hraca
|
||||
rng = Random(21)
|
||||
for _ in range(30):
|
||||
r = Round(rng.randrange(4), rng.randrange(4))
|
||||
players = [RandomPlayer(Random(rng.random())) for _ in range(4)]
|
||||
for _ in range(4):
|
||||
seat = r.get_active_player()
|
||||
r.add_player_guess(seat, players[seat].guess(r, seat))
|
||||
while not r.is_completed():
|
||||
seat = r.get_active_player()
|
||||
r.play_card(seat, index_card(players[seat].play(r, seat)))
|
||||
for other, banned in deduce_voids(r).items():
|
||||
held = {c.color for c in r.player_cards[other]}
|
||||
self.assertFalse(held & banned,
|
||||
f'void {banned} vs ruka {held}')
|
||||
|
||||
def test_deal_consistent_respects_voids(self):
|
||||
rng = Random(22)
|
||||
unknown = [c for c in cards][:20]
|
||||
voids = {1: {Card_colors['HEARTS']}, 2: set(),
|
||||
3: {Card_colors['LEAVES'], Card_colors['BELLS']}}
|
||||
for _ in range(20):
|
||||
hands = deal_consistent(unknown, {1: 4, 2: 4, 3: 4}, voids, rng)
|
||||
self.assertTrue(all(len(h) == 4 for h in hands.values()))
|
||||
for seat, banned in voids.items():
|
||||
self.assertFalse({c.color for c in hands[seat]} & banned)
|
||||
|
||||
|
||||
class McPlayerCase(unittest.TestCase):
|
||||
def test_plays_legal_full_rounds(self):
|
||||
rng = Random(23)
|
||||
env = RoundEnv(rng)
|
||||
players = [McPlayer(Random(24), n_samples=20, play_samples=8),
|
||||
McPlayer(Random(25), n_samples=20, play_samples=8,
|
||||
use_voids=False)] \
|
||||
+ [RandomPlayer(Random(s)) for s in (26, 27)]
|
||||
for round_number in range(8):
|
||||
rewards = play_round(players, env, round_number)
|
||||
self.assertEqual(len(rewards), 4)
|
||||
|
||||
def test_duck_scenario(self):
|
||||
# tip 0, kopku vedie sused LEAVES C8 a ja som HNED na tahu (MC hrac
|
||||
# stavia kopku poctivo, takze na rozdiel od pravidlovej heuristiky
|
||||
# vyzaduje konzistentne poradie): mam ACE aj C7 -> podlezt sedmickou
|
||||
hand0 = [Card(Card_colors['LEAVES'], Card_values['ACE']),
|
||||
Card(Card_colors['LEAVES'], Card_values['C7'])]
|
||||
hand3 = [Card(Card_colors['LEAVES'], Card_values['C8']),
|
||||
Card(Card_colors['LEAVES'], Card_values['C9'])]
|
||||
rest = [c for c in cards if c not in hand0 + hand3]
|
||||
deck = rest[:24] + hand0 + rest[24:26] + rest[26:28] + hand3
|
||||
r = Round(6, 0, deck, shuffler=lambda l: None)
|
||||
r.add_player_guess(0, 0)
|
||||
r.add_player_guess(1, 0)
|
||||
r.add_player_guess(2, 0)
|
||||
r.add_player_guess(3, 1) # najvyssi tip -> vynasa hrac 3, po nom ja
|
||||
r.play_card(3, hand3[0])
|
||||
self.assertEqual(r.get_active_player(), 0)
|
||||
action = McPlayer(Random(28), play_samples=30).play(r, 0)
|
||||
self.assertEqual(index_card(action), hand0[1])
|
||||
|
||||
|
||||
class EvaluateCase(unittest.TestCase):
|
||||
def test_full_random_matchup_runs(self):
|
||||
rng = Random(11)
|
||||
stats = evaluate([RandomPlayer(rng) for _ in range(4)], 40, rng)
|
||||
for s in stats:
|
||||
self.assertEqual(s['rounds'], 40)
|
||||
self.assertGreaterEqual(s['avg_points'], 0)
|
||||
self.assertLessEqual(s['hit_rate'], 1)
|
||||
|
||||
def test_heuristic_beats_random(self):
|
||||
rng = Random(13)
|
||||
players = [HeuristicPlayer(rng, n_samples=40)] \
|
||||
+ [RandomPlayer(rng) for _ in range(3)]
|
||||
stats = evaluate(players, 120, rng)
|
||||
heuristic, randoms = stats[0], stats[1:]
|
||||
best_random = max(s['avg_points'] for s in randoms)
|
||||
self.assertGreater(heuristic['avg_points'], best_random)
|
||||
self.assertGreater(heuristic['hit_rate'],
|
||||
max(s['hit_rate'] for s in randoms))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main(verbosity=2)
|
||||
@@ -0,0 +1,200 @@
|
||||
"""Testy PPO pipeline (rl/model, rl/selfplay, rl/policy_player, rl/train).
|
||||
|
||||
Vyzaduju torch (requirements-rl.txt); bez neho sa cely modul preskoci --
|
||||
ostatne suity (engine, encoding, boti) na torchi nezavisia.
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from random import Random
|
||||
|
||||
try:
|
||||
import torch
|
||||
except ImportError: # pragma: no cover
|
||||
raise unittest.SkipTest('torch nie je nainstalovany (requirements-rl.txt)')
|
||||
|
||||
from bridzik import Round
|
||||
from rl.encoding import (
|
||||
N_GUESS_ACTIONS, N_PLAY_ACTIONS, OBS_DIM,
|
||||
encode_observation, guess_mask, play_mask,
|
||||
)
|
||||
from rl.evaluate import evaluate, play_round
|
||||
from rl.env import RoundEnv
|
||||
from rl.model import BridzikNet, mask_tensor, masked_categorical, obs_tensor
|
||||
from rl.players import RandomPlayer
|
||||
from rl.policy_player import NeuralPlayer
|
||||
from rl.selfplay import _assign_seats, collect_episodes
|
||||
from rl.train import ppo_update
|
||||
|
||||
|
||||
class ModelCase(unittest.TestCase):
|
||||
def test_output_shapes(self):
|
||||
net = BridzikNet(hidden=32)
|
||||
obs = torch.zeros((5, OBS_DIM))
|
||||
guess_logits, play_logits, value = net(obs)
|
||||
self.assertEqual(guess_logits.shape, (5, N_GUESS_ACTIONS))
|
||||
self.assertEqual(play_logits.shape, (5, N_PLAY_ACTIONS))
|
||||
self.assertEqual(value.shape, (5,))
|
||||
|
||||
def test_masked_categorical_never_samples_illegal(self):
|
||||
torch.manual_seed(0)
|
||||
logits = torch.zeros((1, 9))
|
||||
mask = torch.tensor([[False, True, False, True, False,
|
||||
False, False, False, False]])
|
||||
dist = masked_categorical(logits, mask)
|
||||
samples = dist.sample((200,))
|
||||
self.assertTrue(set(samples.flatten().tolist()) <= {1, 3})
|
||||
# entropia a log_prob su konecne aj s -inf logitmi
|
||||
self.assertTrue(torch.isfinite(dist.entropy()).all())
|
||||
self.assertTrue(torch.isfinite(dist.log_prob(torch.tensor([1]))).all())
|
||||
|
||||
def test_encoding_tensors(self):
|
||||
r = Round(3, 0)
|
||||
obs = obs_tensor(encode_observation(r, 0))
|
||||
self.assertEqual(obs.shape, (OBS_DIM,))
|
||||
self.assertEqual(mask_tensor(guess_mask(r)).shape, (N_GUESS_ACTIONS,))
|
||||
self.assertEqual(mask_tensor(play_mask(r, 0)).shape, (N_PLAY_ACTIONS,))
|
||||
|
||||
|
||||
class NeuralPlayerCase(unittest.TestCase):
|
||||
def test_untrained_net_plays_legal_full_rounds(self):
|
||||
torch.manual_seed(1)
|
||||
net = BridzikNet(hidden=32)
|
||||
env = RoundEnv(Random(2))
|
||||
players = [NeuralPlayer(net, greedy=True),
|
||||
NeuralPlayer(net, greedy=False),
|
||||
RandomPlayer(Random(3)), RandomPlayer(Random(4))]
|
||||
# dohratie kola bez BridzikException = vsetky tahy legalne
|
||||
for round_number in range(8):
|
||||
rewards = play_round(players, env, round_number)
|
||||
self.assertEqual(len(rewards), 4)
|
||||
|
||||
|
||||
class SelfPlayCase(unittest.TestCase):
|
||||
def test_collect_episodes_batch_consistency(self):
|
||||
torch.manual_seed(5)
|
||||
net = BridzikNet(hidden=32)
|
||||
batch = collect_episodes(net, n_episodes=6, rng=Random(6))
|
||||
|
||||
n = batch['obs'].shape[0]
|
||||
self.assertGreater(n, 0)
|
||||
for key, width in (('guess_mask', N_GUESS_ACTIONS),
|
||||
('play_mask', N_PLAY_ACTIONS)):
|
||||
self.assertEqual(batch[key].shape, (n, width))
|
||||
for key in ('phase_play', 'action', 'logp', 'value', 'ret'):
|
||||
self.assertEqual(batch[key].shape, (n,))
|
||||
|
||||
# kazda epizoda ma prave 4 guess kroky -> pocet guess krokov = 4*epizody
|
||||
self.assertEqual(int((~batch['phase_play']).sum()), 4 * 6)
|
||||
# return je bud 0 alebo (10+tip)/REWARD_SCALE, cize v (0.55, 1.0]
|
||||
for r in batch['ret'].tolist():
|
||||
self.assertTrue(r == 0.0 or 10.0 / 18.0 <= r <= 1.0)
|
||||
# akcia bola vzdy legalna podla ulozenej masky svojej fazy
|
||||
for i in range(n):
|
||||
mask = batch['play_mask'][i] if batch['phase_play'][i] \
|
||||
else batch['guess_mask'][i]
|
||||
self.assertTrue(bool(mask[batch['action'][i]]))
|
||||
self.assertGreaterEqual(batch['mean_points'], 0.0)
|
||||
|
||||
|
||||
class OpponentMixingCase(unittest.TestCase):
|
||||
def test_assign_seats_always_keeps_a_net_seat(self):
|
||||
rng = Random(20)
|
||||
marker = object()
|
||||
for _ in range(200):
|
||||
seats = _assign_seats(rng, 1.0, 0.0, marker, marker)
|
||||
self.assertIn(None, seats.values()) # aj pri mix_random=1.0
|
||||
self.assertEqual(set(seats), {0, 1, 2, 3})
|
||||
|
||||
def test_mixed_episodes_record_only_net_seats(self):
|
||||
torch.manual_seed(21)
|
||||
net = BridzikNet(hidden=32)
|
||||
# mix_random=1.0 -> presne jedno sietove sedadlo na epizodu
|
||||
batch = collect_episodes(net, n_episodes=5, rng=Random(22),
|
||||
mix_random=1.0)
|
||||
# 1 sietove sedadlo = presne 1 guess krok na epizodu
|
||||
self.assertEqual(int((~batch['phase_play']).sum()), 5)
|
||||
for i in range(batch['obs'].shape[0]):
|
||||
mask = batch['play_mask'][i] if batch['phase_play'][i] \
|
||||
else batch['guess_mask'][i]
|
||||
self.assertTrue(bool(mask[batch['action'][i]]))
|
||||
for r in batch['ret'].tolist():
|
||||
self.assertTrue(r == 0.0 or 10.0 / 18.0 <= r <= 1.0)
|
||||
|
||||
def test_mixed_episodes_with_heuristic(self):
|
||||
torch.manual_seed(23)
|
||||
net = BridzikNet(hidden=32)
|
||||
batch = collect_episodes(net, n_episodes=4, rng=Random(24),
|
||||
mix_heuristic=0.5, heuristic_samples=10)
|
||||
n_guess = int((~batch['phase_play']).sum())
|
||||
self.assertGreaterEqual(n_guess, 4) # aspon 1 sietove sedadlo/epizodu
|
||||
self.assertLessEqual(n_guess, 16)
|
||||
self.assertGreaterEqual(batch['mean_points'], 0.0)
|
||||
|
||||
|
||||
class PpoUpdateCase(unittest.TestCase):
|
||||
def test_update_changes_params_and_is_finite(self):
|
||||
torch.manual_seed(7)
|
||||
net = BridzikNet(hidden=32)
|
||||
optimizer = torch.optim.Adam(net.parameters(), lr=1e-3)
|
||||
batch = collect_episodes(net, n_episodes=8, rng=Random(8))
|
||||
|
||||
before = [p.detach().clone() for p in net.parameters()]
|
||||
stats = ppo_update(net, optimizer, batch, epochs=2, minibatch=64)
|
||||
|
||||
for key in ('policy_loss', 'value_loss', 'entropy'):
|
||||
self.assertTrue(torch.isfinite(torch.tensor(stats[key])))
|
||||
changed = any(
|
||||
not torch.equal(b, a.detach())
|
||||
for b, a in zip(before, net.parameters())
|
||||
)
|
||||
self.assertTrue(changed)
|
||||
|
||||
def test_value_head_learns_constant_reward(self):
|
||||
# sanity uciaceho kroku: na batchi s konstantnym returnom sa value
|
||||
# loss po par updatoch zmensi
|
||||
torch.manual_seed(9)
|
||||
net = BridzikNet(hidden=32)
|
||||
optimizer = torch.optim.Adam(net.parameters(), lr=3e-3)
|
||||
batch = collect_episodes(net, n_episodes=8, rng=Random(10))
|
||||
batch['ret'] = torch.full_like(batch['ret'], 12.0 / 18.0)
|
||||
|
||||
first = ppo_update(net, optimizer, batch, epochs=1, minibatch=4096)
|
||||
for _ in range(10):
|
||||
last = ppo_update(net, optimizer, batch, epochs=1, minibatch=4096)
|
||||
self.assertLess(last['value_loss'], first['value_loss'])
|
||||
|
||||
|
||||
class CheckpointCase(unittest.TestCase):
|
||||
def test_save_load_roundtrip_with_hidden(self):
|
||||
import os
|
||||
import tempfile
|
||||
from rl.train import load_checkpoint, save_checkpoint
|
||||
|
||||
torch.manual_seed(13)
|
||||
net = BridzikNet(hidden=48)
|
||||
path = os.path.join(tempfile.gettempdir(), 'bridzik_ckpt_test.pt')
|
||||
try:
|
||||
save_checkpoint(net, 48, path)
|
||||
loaded = load_checkpoint(path)
|
||||
self.assertEqual(loaded.trunk[0].out_features, 48)
|
||||
obs = torch.zeros((1, OBS_DIM))
|
||||
for a, b in zip(net(obs), loaded(obs)):
|
||||
self.assertTrue(torch.equal(a, b))
|
||||
finally:
|
||||
os.remove(path)
|
||||
|
||||
|
||||
class EvaluateIntegrationCase(unittest.TestCase):
|
||||
def test_neural_player_in_harness(self):
|
||||
torch.manual_seed(11)
|
||||
net = BridzikNet(hidden=32)
|
||||
rng = Random(12)
|
||||
stats = evaluate(
|
||||
[NeuralPlayer(net)] + [RandomPlayer(rng) for _ in range(3)],
|
||||
n_rounds=20, rng=rng,
|
||||
)
|
||||
self.assertEqual(stats[0]['rounds'], 20)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main(verbosity=2)
|
||||
Reference in New Issue
Block a user