"""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)