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): stem Conv2d(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 pons/ce seen during training — here generation 12 (5 min in):

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).

Downloads last month

-

Downloads are not tracked for this model. How to track
Video Preview
loading