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>
50 lines
1.7 KiB
Python
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)
|