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>
This commit is contained in:
+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)
|
||||
Reference in New Issue
Block a user