arena-2.5-mcts-c4
An AlphaZero policy + value network for Connect-4, trained by self-play for the ARENA 3.0 curriculum, chapter [2.5] MCTS & AlphaZero. Small (656 k params), strong, and benchmarked against a perfect solver rather than a heuristic opponent.
- Game: Connect-4 (7 columns × 6 rows).
- Params: 655,560 (2.6 MB, fp32).
- What it predicts: from a board position, a policy (distribution over the 7 columns) and a value (expected game result for the side to move, in [−1, 1]).
Architecture
A small AlphaZero-style ResNet (defined as Connect4Model in ARENA 3.0 chapter 2.5):
- Input:
(B, 3, 6, 7)NCHW, canonicalised to the mover's perspective — channels[empty, mover, opponent](the absolute board is stored[empty, player1, player2]; the two player planes are swapped when it is player-2's turn). - Trunk (
features): stemConv2d(3→128, 3×3, bias=False)+ BatchNorm + ReLU → 2 ResBlocks (128 ch) →(B, 128, 6, 7). - Heads: actor (policy) → logits over 7 columns; critic (value) → scalar in [−1, 1] (tanh), from the mover's perspective.
Training
Self-play AlphaZero with the chapter's fully batched MCTS (BatchedMCTS: root-parallel
flat-tensor trees, one batched network forward per simulation, no host↔device syncs). The network
is trained on the visit-count policy targets and negamax game-outcome value targets. This checkpoint
is the chapter's default recipe (AZConfig() defaults), trained for 32 generations (13 min on
one RTX A4000) with the model checkpointed at the best 5 min in):pons/ce seen during training — here
generation 12 (
| self-play games / generation | 4096 |
| MCTS simulations / move | 16 |
| generations | 32 trained; checkpoint = gen 12 (best solver cross-entropy) |
| replay buffer | last 4 generations |
| root Dirichlet noise | α = 10/7 ≈ 1.43, ε = 0.25 |
| temperature | Ï„ = 1 (visit-count sampling) throughout self-play |
| c_puct | 1.0 |
| LR schedule | cosine 5e-3 → 2e-5 over the run (AdamW, wd 1e-4, grad-clip 1.0) |
| minibatch | 1024 |
| loss | policy cross-entropy + value MSE (equal weight) |
Best-checkpoint selection matters: at this horizon the model over-trains — solver cross-entropy bottoms out early (0.444 at gen 12) and then degrades ~15% by gen 32 (0.513), so the published weights are the gen-12 snapshot, not the end of training.
Training curves (per-generation loss / lr / Pons metrics, logged live):
wandb run ual3fkoc.
Why these choices (from the ARENA 2.5 full-run sweeps):
- Dirichlet root noise is load-bearing — without it the policy collapses (self-play degenerates to one drawish line, the value head dies, strength craters). It is the exploration floor.
- c_puct 1.5→1.0, lr 1e-3→5e-3 and buffer 8→4 generations were tuned by a full-run sweep (pons CE ~0.466→0.420 at ~¼ the compute); every gain came from fresher self-play data — capacity / loss-weighting / regularisation changes were inert. The usable lr band is ~3e-3..6e-3 (≥7e-3 is seed-unstable).
- Matching the cosine-LR horizon to the run length avoids the over-training drift seen in long runs (the model peaks early then degrades at near-max LR) — so this run anneals into convergence.
Evaluation vs a perfect solver ("God's eval")
Rather than win-rate against a weak heuristic, this model is scored against Pascal Pons' perfect Connect-4 solver on a frozen set of 6,705 decisive positions (positions where the mover can reach ≥2 distinct outcome classes of win/draw/loss — i.e. the move actually matters), spanning plies 2–36. For each position we compare the raw policy (no search) to the solver's set of game-theoretically optimal moves:
| metric (raw policy, no MCTS) | value |
|---|---|
| optimal-move accuracy (argmax ∈ optimal set) | 85.0 % |
| blunder rate (argmax is sub-optimal) | 15.0 % |
| cross-entropy to the optimal set, −log Σ p(optimal) | 0.44 |
| value-head sign accuracy (W/L) | 86.8 % |
By game phase (optimal-move accuracy): opening 79.8 %, midgame 87.9 %, endgame 87.3 %.
These numbers were reproduced from a fresh hub download with the chapter's eval_hf_c4.py
(pons/acc = 0.8501, pons/ce = 0.4440 on the frozen set). With MCTS at decision time the agent
is considerably stronger than the raw policy numbers above.
Usage
import torch
# Connect4Model is defined in ARENA 3.0, chapter2_rl exercises, part5_mcts_alphazero (solutions.py).
from solutions import Connect4Model # or copy the class from the chapter
from huggingface_hub import hf_hub_download
device = "cuda" if torch.cuda.is_available() else "cpu"
ckpt = hf_hub_download("davidquarel/arena-2.5-mcts-c4", "arena-2.5-mcts-c4.pt")
model = Connect4Model(device)
model.load_state_dict(torch.load(ckpt, map_location=device))
model.eval()
# obs: (B, 3, 6, 7) absolute board [empty, player1, player2]; canonicalise to the mover before calling.
# value, logits = model(canonicalise_obs(obs, is_player1)) # value (B,), logits (B,7)
For best play, wrap the network in MCTS (the chapter's BatchedMCTS) at decision time rather than
taking the raw policy argmax.
Caveats
- Connect-4 is small and solved, so a strong agent is near-perfect and has few exploitable holes; the raw policy is ~86 % optimal and a perfect solver beats it. The methods (perfect-solver eval, adversarial probing) transfer to larger games where they're more dramatic.
- This is an educational artifact for the ARENA curriculum, not a SOTA game engine. The default recipe trades a few points of accuracy for a much shorter run; the chapter's stronger recipe (64 sims/move, longer horizon) reaches ~88 % on the same eval.
- Checkpoint selection optimises cross-entropy to the optimal-move set (tie-invariant, the chapter's headline metric); argmax accuracy peaks at a slightly different point in training.
Provenance
Trained 2026-06-12 with train_c4_best.py (the chapter's rewritten batched MCTS code,
default AZConfig recipe, best-pons/ce checkpointing) in ARENA 3.0 (davidquarel/ARENA_3.0,
chapter2_rl/.../part5_mcts_alphazero). Training curves:
wandb. Evaluated with
pascal_pons/eval_pons.py; the perfect-solver evaluation harness builds a frozen, md5-verified
dataset from Pascal Pons' connect4 solver
(hosted at davidquarel/connect4-pons-eval).