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:
+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),
|
||||
}
|
||||
Reference in New Issue
Block a user