Files
timandClaude Fable 5 8f2449a408 RL: self-play PPO trening siete
BridzikNet (trup + guess/play/value hlavy s maskovanim), self-play
generator so zdielanou sietou na 4 sedadlach a opponent mixingom
(random/heuristicke sedadla pre robustnost), vlastny clipped-PPO
so skalovanim odmien a lr/entropy annealom. Torch je len trenovacia
zavislost na hoste (requirements-rl.txt); checkpointy a logy su
gitignorovane. Spustenie: py -m rl.train.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-07 18:49:55 +02:00

50 lines
1.7 KiB
Python

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