general-bots-training/scripts/train.py
2026-08-07 22:06:27 +02:00

179 lines
6.1 KiB
Python

"""Train a Generals.io policy with PPO against a configurable static opponent.
Run with `.venv/bin/python scripts/train.py` or `uv run python scripts/train.py`.
Override defaults with OmegaConf arguments such as `num_envs=32 lr=1e-4`.
"""
import time
from dataclasses import dataclass
import equinox as eqx
import jax
import jax.numpy as jnp
import jax.random as jrandom
import optax
from generals.core.env import GeneralsEnv
from omegaconf import DictConfig, OmegaConf
from general_bots_training.network import PolicyValueNetwork
from general_bots_training.opponents import make_opponent
from general_bots_training.ppo import compute_advantages_and_returns, make_train_epoch
from general_bots_training.rollout import make_collect_rollout
@dataclass
class TrainingConfig:
grid_dims: tuple[int, int] = (21, 21)
truncation: int = 500
num_envs: int = 256
rollout_steps: int = 256
num_iterations: int = 500
num_epochs: int = 1
minibatch_size: int = 256
lr: float = 3e-4
gamma: float = 0.99
lam: float = 0.95
clip: float = 0.2
value_coef: float = 0.5
entropy_coef: float = 0.01
log_every: int = 10
checkpoint_path: str = "ppo_model.eqx"
resume_from: str | None = None
opponent: str = "random"
seed: int = 0
def load_config(args: list[str] | None = None) -> DictConfig:
"""Merge CLI overrides into the typed default training configuration."""
defaults = OmegaConf.structured(TrainingConfig)
return OmegaConf.merge(defaults, OmegaConf.from_cli(args)) # ty: ignore
def initialize_network(key, checkpoint_path: str | None = None):
"""Initialize a policy network, optionally restoring serialized leaves."""
network = PolicyValueNetwork(key, in_channels=14)
if checkpoint_path is not None:
network = eqx.tree_deserialise_leaves(checkpoint_path, network)
return network
def main(config: DictConfig):
key = jrandom.PRNGKey(config.seed)
key, net_key, pool_key = jrandom.split(key, 3)
opponent = make_opponent(config.opponent)
network = initialize_network(net_key, config.resume_from)
env = GeneralsEnv(grid_dims=tuple(config.grid_dims), truncation=config.truncation)
pool, _ = env.reset(pool_key)
optimizer = optax.adam(config.lr)
opt_state = optimizer.init(eqx.filter(network, eqx.is_array))
params, _ = eqx.partition(network, eqx.is_array)
n_params = sum(x.size for x in jax.tree.leaves(params))
print("Generals.io PPO (fog)")
print(OmegaConf.to_yaml(config, resolve=True).rstrip())
print(f"device: {jax.devices()[0]}")
print(f"params: {n_params:,}")
if config.resume_from is not None:
print(f"resumed model weights from {config.resume_from}")
print()
collect = make_collect_rollout(env, config.rollout_steps, opponent)
train_epoch = make_train_epoch(
optimizer,
config.minibatch_size,
clip=config.clip,
value_coef=config.value_coef,
entropy_coef=config.entropy_coef,
)
# Initial per-env states.
key, init_key = jrandom.split(key)
init_keys = jrandom.split(init_key, config.num_envs)
states = jax.vmap(env.init_state)(init_keys)
states = states._replace(
pool_idx=jnp.arange(config.num_envs, dtype=states.pool_idx.dtype) % env.pool_size
)
# Compile with disposable outputs so warmup does not advance training state.
print("warming up (jit compile)...")
warm_states, transitions, (warm_key, last_next_obs) = collect(states, pool, network, key)
next_value = jax.vmap(network.value)(last_next_obs)
advantages, returns = compute_advantages_and_returns(
transitions["reward"],
transitions["value"],
next_value,
transitions["done"],
gamma=config.gamma,
lam=config.lam,
)
batch = (
transitions["obs"],
transitions["mask"],
transitions["action"],
transitions["logprob"],
advantages,
returns,
)
warm_key, epoch_key = jrandom.split(warm_key)
warm_network, _, _ = train_epoch(network, opt_state, batch, epoch_key)
jax.block_until_ready((warm_states, warm_network))
print("warming up done\n")
print("training...")
for it in range(config.num_iterations):
t0 = time.time()
states, transitions, (key, last_next_obs) = collect(states, pool, network, key)
next_value = jax.vmap(network.value)(last_next_obs)
advantages, returns = compute_advantages_and_returns(
transitions["reward"],
transitions["value"],
next_value,
transitions["done"],
gamma=config.gamma,
lam=config.lam,
)
batch = (
transitions["obs"],
transitions["mask"],
transitions["action"],
transitions["logprob"],
advantages,
returns,
)
epoch_losses = []
for _ in range(config.num_epochs):
key, epoch_key = jrandom.split(key)
network, opt_state, loss = train_epoch(network, opt_state, batch, epoch_key)
epoch_losses.append(loss)
jax.block_until_ready(network)
loss = jnp.mean(jnp.stack(epoch_losses))
elapsed = time.time() - t0
if it % config.log_every == 0:
player_zero = transitions["player"] == 0
dones = transitions["done"] & player_zero
winner = transitions["winner"]
num_episodes = int(dones.sum())
wins = int(jnp.sum(dones & (winner == 0)))
losses_count = int(jnp.sum(dones & (winner == 1)))
win_rate = wins / max(num_episodes, 1) * 100
sps = (config.num_envs * config.rollout_steps) / elapsed
print(
f"iter {it:4d} | loss {float(loss):.4f} | "
f"reward {float(transitions['reward'].mean()):+.4f} | "
f"eps {num_episodes:3d} | wins {wins:2d}/{num_episodes} "
f"({win_rate:.0f}%) | losses {losses_count:2d} | "
f"sps {sps:7.0f} | {elapsed:.2f}s"
)
eqx.tree_serialise_leaves(config.checkpoint_path, network)
print(f"\nmodel saved to {config.checkpoint_path}")
if __name__ == "__main__":
main(load_config())