No description
  • Python 98.9%
  • Shell 1%
  • Just 0.1%
Find a file
Repository files (latest commit first)
Filename Latest commit message Latest commit date
2026-08-08 21:25:35 +02:00
agents/mcts checkpoint 2026-08-07 22:06:27 +02:00
competition checkpoint 2026-08-08 21:25:35 +02:00
scripts checkpoint 2026-08-08 21:25:35 +02:00
src/general_bots_training added training against checkpoint 2026-08-08 01:07:26 +02:00
tests added training against checkpoint 2026-08-08 01:07:26 +02:00
.gitignore checkpoint 2026-08-08 21:25:35 +02:00
.python-version checkpoint 2026-08-07 22:06:27 +02:00
competition-submission.zip checkpoint 2026-08-08 21:25:35 +02:00
competition-submission.zip.bak checkpoint 2026-08-08 21:25:35 +02:00
Justfile checkpoint 2026-08-08 21:25:35 +02:00
ppo_model.eqx checkpoint 2026-08-08 21:25:35 +02:00
ppo_model.eqx.1 checkpoint 2026-08-07 22:06:27 +02:00
ppo_model.eqx.2 checkpoint 2026-08-07 22:06:27 +02:00
ppo_model.eqx.3 added evaluation script 2026-08-08 01:06:37 +02:00
ppo_model.eqx.4 added evaluation script 2026-08-08 01:06:37 +02:00
ppo_model.eqx.5 added evaluation script 2026-08-08 01:06:37 +02:00
ppo_model.eqx.6 added training against checkpoint 2026-08-08 01:07:26 +02:00
ppo_model.eqx.7 checkpoint 2026-08-08 21:25:35 +02:00
ppo_model_h100.eqx checkpoint 2026-08-08 21:25:35 +02:00
ppo_model_h100.eqx.1 added evaluation script 2026-08-08 01:06:37 +02:00
ppo_model_h100.eqx.2 checkpoint 2026-08-08 21:25:35 +02:00
ppo_model_h100.eqx.3 checkpoint 2026-08-08 21:25:35 +02:00
ppo_model_h100.eqx.4 checkpoint 2026-08-08 21:25:35 +02:00
pyproject.toml added evaluation script 2026-08-08 01:06:37 +02:00
README.md checkpoint 2026-08-07 22:06:27 +02:00
requirements-competition.txt checkpoint 2026-08-08 21:25:35 +02:00
requirements-vendor.txt checkpoint 2026-08-08 21:25:35 +02:00
requirements.txt added evaluation script 2026-08-08 01:06:37 +02:00
uv.lock added evaluation script 2026-08-08 01:06:37 +02:00

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+1 logits. Invalid moves are masked to 1e9 via compute_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:

  1. Observe both players from the current state.
  2. Sample p0's action from the policy network. For a static opponent, obtain p1's action from RandomAgent, ExpanderAgent, or HunterAgent; in self_play, sample p1 independently from the same current network.
  3. Step the env (vmapped), which auto-resets from the pool on done.
  4. Compute the shaped reward for p0 with composite_reward_fn from the pre-step and post-step observations. The post-step observation is taken from timestep.last_state (the state before auto-reset) so terminal and shaping rewards are computed against the actual end-of-episode board.
  5. 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 reverse lax.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 with eqx.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 with eqx.filter_grad; eqx.filter_vmap on 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 ~2044% 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 generals package is pinned via git in [tool.uv.sources]; the generals-bots/ subdir is a clone for reference and is not part of the build.
  • GPU isn't visible from the Zed sandbox (cuInit fails → CPU fallback). Run scripts/train.py from your local machine or a GPU host for CUDA.

Roadmap

Planned extensions, in rough priority order:

  1. Opponent league — extend current-policy self-play with frozen historical snapshots to reduce strategy collapse.
  2. Curriculum — step up from 4×4 to larger grids, then to GeneralsEnv(mode="competition") (variable 1821 grids, 1200-step truncation, build_castles + deathtouch modifiers).
  3. Algorithm swap — the PPO logic is isolated in ppo.py; REINFORCE or another algorithm can replace it without touching the rollout or network.
  4. Evaluation harness — match the trained policy against the bundled ExpanderAgent and the competition's stdio bots via competition/matchup.py.