- Python 98.9%
- Shell 1%
- Just 0.1%
| Filename | Latest commit message | Latest commit date |
|---|---|---|
| agents/mcts | ||
| scripts | ||
| src/general_bots_training | ||
| tests | ||
| .python-version | ||
| ppo_model.eqx | ||
| ppo_model.eqx.1 | ||
| ppo_model.eqx.2 | ||
| ppo_model.eqx.3 | ||
| ppo_model.eqx.4 | ||
| ppo_model.eqx.5 | ||
| ppo_model_h100.eqx | ||
| ppo_model_h100.eqx.1 | ||
| pyproject.toml | ||
| README.md | ||
| requirements.txt | ||
| uv.lock | ||
general-bots-training
Reinforcement learning training for the generals.bot competition. Trains a policy with PPO against the generals JAX environment.
Status
v1 — minimal working PPO loop. 4×4 grid, fog of war, composite reward shaping, configurable static opponents, and two-sided self-play. The code is structured so planned extensions such as opponent leagues, curriculum, and the competition ruleset fit without rewriting the core.
Quick start
# install (editable, picks up src/general_bots_training)
uv sync
# train (uses GPU if available, falls back to CPU)
uv run python scripts/train.py
In a network-restricted sandbox, uv run may fail to re-resolve the build backend; use the venv directly:
.venv/bin/python scripts/train.py
Checkpoints are written to ppo_model.eqx by default. Override this with checkpoint_path=....
Configuration
scripts/train.py defines a typed OmegaConf configuration. Override defaults with key=value arguments; CLI values are merged over the structured defaults:
uv run python scripts/train.py num_envs=32 rollout_steps=64 lr=1e-4
uv run python scripts/train.py grid_dims=[6,6] opponent=expander
uv run python scripts/train.py opponent=hunter checkpoint_path=models/ppo-hunter.eqx
uv run python scripts/train.py resume_from=models/ppo-hunter.eqx opponent=expander
uv run python scripts/train.py opponent=self_play
Unknown keys and incompatible value types are rejected.
| key | default | notes |
|---|---|---|
grid_dims |
[4, 4] |
start small; curriculum to larger grids later |
truncation |
500 |
max turns before a game is scored as a draw |
num_envs |
256 |
parallel games; tune to GPU VRAM |
rollout_steps |
256 |
steps per rollout before a PPO update |
num_iterations |
500 |
PPO update count |
num_epochs |
1 |
epochs over each rollout buffer |
minibatch_size |
256 |
|
lr |
3e-4 |
Adam |
gamma / lam |
0.99 / 0.95 |
GAE |
clip |
0.2 |
PPO ratio clip |
value_coef / entropy_coef |
0.5 / 0.01 |
|
log_every |
10 |
iterations between progress logs |
checkpoint_path |
ppo_model.eqx |
output model path |
resume_from |
null |
optional model checkpoint to continue from |
opponent |
random |
random, expander, hunter, or self_play |
seed |
0 |
JAX random seed |
resume_from restores the policy/value network weights. Existing checkpoints do not contain optimizer state, so Adam starts with fresh moments and the configured learning rate.
Layout
src/general_bots_training/
network.py # equinox conv policy-value net + observation encoding
mcts.py # competition observation adapter and particle PUCT search
ppo.py # reusable GAE, clipped PPO loss, and optimizer helpers
opponents.py # opponent interfaces and named strategy selection
rollout.py # jitted static-opponent and two-sided self-play collection
scripts/
train.py # executable config, training loop, logging, checkpointing
mcts_agent.py # competition stdio inference entrypoint
agents/mcts/
run.sh # local matchup wrapper for ppo_model.eqx
Architecture
Network (network.py)
PolicyValueNetwork is an equinox module:
- Backbone: 4 conv layers (3×3, padding=1) over a 14-channel normalized observation. Armies, army-counts, and timestep are log-normalized; scalar values are broadcast to spatial planes so a plain conv stack can consume them.
- Policy head: 1×1 conv to 9 channels = 4 full-move directions + 4 half-move (split) directions + a spatial pass score. Move channels are flattened and the pass scores are spatially pooled into one global pass action, yielding
8*H*W+1logits. Invalid moves are masked to −1e9 viacompute_valid_move_mask; pass is always available. - Value head: 1×1 conv → global average pool → 2-layer MLP → scalar. Global pooling makes the network grid-size-agnostic, so the same architecture extends to larger boards without reshaping linear layers.
obs_to_tensor encodes a generals.Observation into the (14, H, W) float32 input.
Rollout (rollout.py)
make_collect_rollout(env, num_steps, opponent) accepts a stateless JAX-compatible agent or the self-play marker and returns a jitted function (states, pool, network, key) -> (states, transitions, (key, last_next_obs)). Each step:
- Observe both players from the current state.
- Sample p0's action from the policy network. For a static opponent, obtain p1's action from
RandomAgent,ExpanderAgent, orHunterAgent; inself_play, sample p1 independently from the same current network. - Step the env (vmapped), which auto-resets from the pool on done.
- Compute the shaped reward for p0 with
composite_reward_fnfrom the pre-step and post-step observations. The post-step observation is taken fromtimestep.last_state(the state before auto-reset) so terminal and shaping rewards are computed against the actual end-of-episode board. - Record
(obs, mask, action, logprob, value, reward, done, winner).
The bootstrap observation for the critic is threaded through the lax.scan carry; only the final step's is returned (as last_next_obs) to avoid storing T copies. Static-opponent rollouts produce N trajectories per step. Self-play produces 2N, with observations, actions, shaped rewards, values, and log-probabilities from both player perspectives included in the same PPO update.
PPO (ppo.py)
compute_gae: GAE via reverselax.scan, bootstrapping from the critic value of the post-rollout state (zeroed on done steps).ppo_loss: clipped surrogate + value loss + entropy bonus.make_train_epoch: flattens(T, N)→(T*N), shuffles, minibatches witheqx.filter_grad.
Training loop (scripts/train.py)
Non-mutating warmup (compile) → per iteration: collect rollout → GAE → compute returns from raw advantages → normalize policy advantages → PPO update → log (loss, reward, episodes, win/loss, SPS) → checkpoint at the end.
Competition PUCT (mcts.py)
Run the local stdio bot directly through the bundled matchup driver:
PYTHONPATH=src:generals-bots .venv/bin/python generals-bots/competition/matchup.py \
agents/mcts/run.sh \
generals-bots/competition/agents/expander_python/run.sh \
--mode competition
The bot performs deadline-bounded root PUCT using only the perspective-relative wire observation. Each simulation samples a hidden-state determinization consistent with visible ownership and global opponent totals, samples a simultaneous opponent action from the same policy, applies build-castles and deathtouch transitions, and evaluates the resulting leaf with the critic. Network move/pass logits provide priors; affordable build actions are added with exact legality and heuristic priors so existing checkpoints remain compatible.
The handshake warmup compiles all board-shape-dependent paths before the first action. On a pinned Ryzen 5800X core, a 21×21 search configured for 125 ms completed in approximately 111 ms with seven depth-2 simulations. Results depend on CPU and position complexity.
This is a conservative first particle search, not full information-set MCTS: particles are regenerated from each current observation and do not yet maintain a persistent history belief. Also, the published competition environment manifest includes JAX but not Equinox or generals-bots; agents/mcts/run.sh is therefore a local evaluation wrapper. A submitted bot must bundle those dependencies or export the network/simulator to the sandbox's available runtime.
Key correctness choices
These differ from the experimental reference in generals-bots/examples/_experimental/ppo/:
- Bootstrap GAE from the post-step critic value, not 0. Done steps are zeroed via the done mask, so a fresh reset state's value doesn't contaminate the advantage.
- Post-step observation from
timestep.last_state(pre-auto-reset) so terminal/shaping rewards are correct. The env's auto-reset overwrites the state with a fresh board; using that for reward shaping would attribute the reset board's counts to the just-finished episode. - Thread the pool explicitly through
env.step(vmapped) rather than capturing it as a constant, so it isn't baked into the JIT trace. jax.vmap(network, in_axes=(0,0,None,0))for batched forward — the network is the vmapped callable, so its weight leaves are batched alongside the data. This composes correctly witheqx.filter_grad;eqx.filter_vmapon a closure capturing the network does not.
Validation
Validated end-to-end on CPU (the sandbox has no GPU):
- Compiles in ~20s, ~320 SPS on 32 envs / 200-step rollouts.
- Episodes complete, win/loss counting works, checkpoints save.
- An untrained network wins ~20–44% vs random (random also wins some by accident) — a sensible starting point.
On a 4080 / rented GPU, throughput should be substantially higher; tune NUM_ENVS and ROLLOUT_STEPS to VRAM.
Notes
- The
generalspackage is pinned via git in[tool.uv.sources]; thegenerals-bots/subdir is a clone for reference and is not part of the build. - GPU isn't visible from the Zed sandbox (
cuInitfails → CPU fallback). Runscripts/train.pyfrom your local machine or a GPU host for CUDA.
Roadmap
Planned extensions, in rough priority order:
- Opponent league — extend current-policy self-play with frozen historical snapshots to reduce strategy collapse.
- Curriculum — step up from 4×4 to larger grids, then to
GeneralsEnv(mode="competition")(variable 18–21 grids, 1200-step truncation,build_castles+deathtouchmodifiers). - Algorithm swap — the PPO logic is isolated in
ppo.py; REINFORCE or another algorithm can replace it without touching the rollout or network. - Evaluation harness — match the trained policy against the bundled
ExpanderAgentand the competition's stdio bots viacompetition/matchup.py.