checkpoint
This commit is contained in:
parent
e815d3f9b2
commit
a69868e950
23 changed files with 1269 additions and 0 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -1 +1,2 @@
|
||||||
__pycache__/
|
__pycache__/
|
||||||
|
competition/.venv-competition/
|
||||||
|
|
|
||||||
2
Justfile
Normal file
2
Justfile
Normal file
|
|
@ -0,0 +1,2 @@
|
||||||
|
vendor:
|
||||||
|
pip download -r competition/requirements-vendor.txt -d competition/wheelhouse/ --python-version 3.12 --platform manylinux --no-deps
|
||||||
BIN
competition-submission.zip
Normal file
BIN
competition-submission.zip
Normal file
Binary file not shown.
BIN
competition-submission.zip.bak
Normal file
BIN
competition-submission.zip.bak
Normal file
Binary file not shown.
7
competition/build.sh
Executable file
7
competition/build.sh
Executable file
|
|
@ -0,0 +1,7 @@
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
# One-time intake step: install the vendored wheels offline.
|
||||||
|
# The sandbox has no network, so everything must come from wheelhouse/.
|
||||||
|
set -euo pipefail
|
||||||
|
cd "$(dirname "$0")"
|
||||||
|
python -m pip install --no-index --find-links=wheelhouse -r requirements-vendor.txt
|
||||||
|
echo "[build] vendored wheels installed" >&2
|
||||||
336
competition/competition.py
Normal file
336
competition/competition.py
Normal file
|
|
@ -0,0 +1,336 @@
|
||||||
|
"""Competition stdio agent: direct policy sampling over a trained equinox checkpoint.
|
||||||
|
|
||||||
|
Self-contained competition bot. It speaks the wire protocol in
|
||||||
|
`generals-bots/competition/protocol.py` (handshake, then one observation frame
|
||||||
|
per turn, one action line per turn, EOF on stdin = game over) and, on each
|
||||||
|
turn, samples an action from the policy network exactly as the model is used
|
||||||
|
during training/evaluation — no tree search.
|
||||||
|
|
||||||
|
Unlike `scripts/mcts_agent.py`, this module does **not** import the
|
||||||
|
`general_bots_training` training package. It inlines the policy-value network
|
||||||
|
and the observation encoding so the script can be dropped into the competition
|
||||||
|
sandbox with just the engine (`generals`), `equinox`, `jax`, and `numpy`
|
||||||
|
available.
|
||||||
|
|
||||||
|
This is the search-free counterpart to `scripts/competition_puct.py`: it shares
|
||||||
|
the same network and observation encoding, but selects an action by sampling
|
||||||
|
from the masked policy (as in `PolicyValueNetwork.__call__`) instead of running
|
||||||
|
particle PUCT.
|
||||||
|
|
||||||
|
Run it directly through the bundled matchup driver, e.g.:
|
||||||
|
|
||||||
|
PYTHONPATH=src:generals-bots .venv/bin/python \
|
||||||
|
generals-bots/competition/matchup.py \
|
||||||
|
scripts/competition.py --checkpoint ppo_model.eqx \
|
||||||
|
generals-bots/competition/agents/expander_python/run.sh \
|
||||||
|
--mode competition
|
||||||
|
|
||||||
|
or wrap it in a `run.sh` that passes the checkpoint path.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
|
||||||
|
import equinox as eqx
|
||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import jax.random as jrandom
|
||||||
|
import numpy as np
|
||||||
|
from generals.core import game
|
||||||
|
from generals.core.action import compute_valid_move_mask
|
||||||
|
from generals.core.game import GameState
|
||||||
|
from generals.core.observation import Observation
|
||||||
|
from generals.modifiers import build_castles, deathtouch
|
||||||
|
|
||||||
|
os.environ.setdefault("JAX_PLATFORMS", "cpu")
|
||||||
|
|
||||||
|
PASS_ACTION = np.array([1, 0, 0, 0, 0], dtype=np.int32)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Policy-value network (inlined from general_bots_training.network)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def obs_to_tensor(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Encode an Observation into a (C, H, W) float32 tensor for the network.
|
||||||
|
|
||||||
|
Armies and army-counts are log-normalized; land counts are divided by the
|
||||||
|
number of cells; scalar values are broadcast to spatial planes so a plain
|
||||||
|
conv stack can consume them. Grid-size agnostic.
|
||||||
|
"""
|
||||||
|
H, W = obs.armies.shape
|
||||||
|
armies = jnp.log1p(obs.armies.astype(jnp.float32)) / jnp.log(50.0)
|
||||||
|
|
||||||
|
def bcast(scalar):
|
||||||
|
return jnp.broadcast_to(scalar.astype(jnp.float32), (H, W))
|
||||||
|
|
||||||
|
own_land = bcast(obs.owned_land_count / (H * W))
|
||||||
|
own_army = bcast(jnp.log1p(obs.owned_army_count.astype(jnp.float32)) / jnp.log(50.0))
|
||||||
|
opp_land = bcast(obs.opponent_land_count / (H * W))
|
||||||
|
opp_army = bcast(jnp.log1p(obs.opponent_army_count.astype(jnp.float32)) / jnp.log(50.0))
|
||||||
|
timestep = bcast(jnp.log1p(obs.timestep.astype(jnp.float32)) / jnp.log(1201.0))
|
||||||
|
|
||||||
|
return jnp.stack(
|
||||||
|
[
|
||||||
|
armies,
|
||||||
|
obs.generals.astype(jnp.float32),
|
||||||
|
obs.castles.astype(jnp.float32),
|
||||||
|
obs.mountains.astype(jnp.float32),
|
||||||
|
obs.neutral_cells.astype(jnp.float32),
|
||||||
|
obs.owned_cells.astype(jnp.float32),
|
||||||
|
obs.opponent_cells.astype(jnp.float32),
|
||||||
|
obs.fog_cells.astype(jnp.float32),
|
||||||
|
obs.structures_in_fog.astype(jnp.float32),
|
||||||
|
own_land,
|
||||||
|
own_army,
|
||||||
|
opp_land,
|
||||||
|
opp_army,
|
||||||
|
timestep,
|
||||||
|
],
|
||||||
|
axis=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PolicyValueNetwork(eqx.Module):
|
||||||
|
"""Conv policy-value network.
|
||||||
|
|
||||||
|
Action layout: 4 full-move directions and 4 half-move (split) directions
|
||||||
|
per source cell, plus one global pass action. Invalid moves are masked to
|
||||||
|
-1e9; pass is always available.
|
||||||
|
|
||||||
|
The value head uses global average pooling over the spatial grid, so the
|
||||||
|
network works for any board size without reshaping linear layers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
conv1: eqx.nn.Conv2d
|
||||||
|
conv2: eqx.nn.Conv2d
|
||||||
|
conv3: eqx.nn.Conv2d
|
||||||
|
conv4: eqx.nn.Conv2d
|
||||||
|
policy_conv: eqx.nn.Conv2d
|
||||||
|
value_conv: eqx.nn.Conv2d
|
||||||
|
value_linear1: eqx.nn.Linear
|
||||||
|
value_linear2: eqx.nn.Linear
|
||||||
|
|
||||||
|
def __init__(self, key, in_channels: int = 14, channels=(32, 32, 32, 16)):
|
||||||
|
keys = jrandom.split(key, 8)
|
||||||
|
self.conv1 = eqx.nn.Conv2d(in_channels, channels[0], kernel_size=3, padding=1, key=keys[0])
|
||||||
|
self.conv2 = eqx.nn.Conv2d(channels[0], channels[1], kernel_size=3, padding=1, key=keys[1])
|
||||||
|
self.conv3 = eqx.nn.Conv2d(channels[1], channels[2], kernel_size=3, padding=1, key=keys[2])
|
||||||
|
self.conv4 = eqx.nn.Conv2d(channels[2], channels[3], kernel_size=3, padding=1, key=keys[3])
|
||||||
|
# 9 = 4 dirs (full) + 4 dirs (half) + 1 pass
|
||||||
|
self.policy_conv = eqx.nn.Conv2d(channels[3], 9, kernel_size=1, key=keys[4])
|
||||||
|
self.value_conv = eqx.nn.Conv2d(channels[3], 4, kernel_size=1, key=keys[5])
|
||||||
|
self.value_linear1 = eqx.nn.Linear(4, 64, key=keys[6])
|
||||||
|
self.value_linear2 = eqx.nn.Linear(64, 1, key=keys[7])
|
||||||
|
|
||||||
|
def _features(self, obs):
|
||||||
|
x = jax.nn.relu(self.conv1(obs))
|
||||||
|
x = jax.nn.relu(self.conv2(x))
|
||||||
|
x = jax.nn.relu(self.conv3(x))
|
||||||
|
x = jax.nn.relu(self.conv4(x))
|
||||||
|
return x
|
||||||
|
|
||||||
|
def _value_from_features(self, feat):
|
||||||
|
v = jax.nn.relu(self.value_conv(feat)) # (4, H, W)
|
||||||
|
v = v.mean(axis=(1, 2)) # (4,) global average pool
|
||||||
|
v = jax.nn.relu(self.value_linear1(v))
|
||||||
|
return self.value_linear2(v)[0]
|
||||||
|
|
||||||
|
def policy_value(self, obs, mask):
|
||||||
|
"""Return deterministic masked policy logits and the critic value."""
|
||||||
|
features = self._features(obs)
|
||||||
|
return self._policy_logits(features, mask), self._value_from_features(features)
|
||||||
|
|
||||||
|
def _policy_logits(self, feat, mask):
|
||||||
|
"""Return 8*H*W move logits plus one global pass logit."""
|
||||||
|
logits = self.policy_conv(feat) # (9, H, W)
|
||||||
|
mask_t = jnp.transpose(mask, (2, 0, 1)) # (4, H, W)
|
||||||
|
penalty = (1.0 - mask_t) * -1e9
|
||||||
|
move_penalty = jnp.concatenate([penalty, penalty], axis=0)
|
||||||
|
move_logits = (logits[:8] + move_penalty).reshape(-1)
|
||||||
|
pass_logit = jnp.mean(logits[8])[None]
|
||||||
|
return jnp.concatenate([move_logits, pass_logit])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Action selection (matches training/evaluation rollout)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def competition_move_mask(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Legal move mask, identical to the one used during training/evaluation."""
|
||||||
|
return compute_valid_move_mask(obs.armies, obs.owned_cells, obs.mountains)
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_policy_index(index: int, height: int, width: int) -> np.ndarray:
|
||||||
|
cells = height * width
|
||||||
|
if index == 8 * cells:
|
||||||
|
return PASS_ACTION.copy()
|
||||||
|
encoded_direction, position = divmod(index, cells)
|
||||||
|
row, col = divmod(position, width)
|
||||||
|
split = int(encoded_direction >= 4)
|
||||||
|
direction = encoded_direction - 4 if split else encoded_direction
|
||||||
|
return np.array([0, row, col, direction, split], dtype=np.int32)
|
||||||
|
|
||||||
|
|
||||||
|
def observation_from_wire(
|
||||||
|
timestep: int,
|
||||||
|
owned_land_count: int,
|
||||||
|
owned_army_count: int,
|
||||||
|
opponent_land_count: int,
|
||||||
|
opponent_army_count: int,
|
||||||
|
type_grid,
|
||||||
|
owner_grid,
|
||||||
|
army_grid,
|
||||||
|
) -> Observation:
|
||||||
|
"""Convert a competition wire frame into the network's Observation type."""
|
||||||
|
types = jnp.asarray(type_grid, dtype=jnp.int32)
|
||||||
|
owners = jnp.asarray(owner_grid, dtype=jnp.int32)
|
||||||
|
armies = jnp.asarray(army_grid, dtype=jnp.int32)
|
||||||
|
visible = (types != 0) & (types != 5)
|
||||||
|
return Observation(
|
||||||
|
armies=armies,
|
||||||
|
generals=types == 4,
|
||||||
|
castles=types == 3,
|
||||||
|
mountains=types == 2,
|
||||||
|
neutral_cells=visible & (owners == 0) & (types != 2),
|
||||||
|
owned_cells=owners == 1,
|
||||||
|
opponent_cells=owners == 2,
|
||||||
|
fog_cells=types == 0,
|
||||||
|
structures_in_fog=types == 5,
|
||||||
|
owned_land_count=jnp.int32(owned_land_count),
|
||||||
|
owned_army_count=jnp.int32(owned_army_count),
|
||||||
|
opponent_land_count=jnp.int32(opponent_land_count),
|
||||||
|
opponent_army_count=jnp.int32(opponent_army_count),
|
||||||
|
timestep=jnp.int32(timestep),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def competition_step(state: GameState, actions: jnp.ndarray):
|
||||||
|
state, actions = build_castles.apply_build_actions(state, actions)
|
||||||
|
return deathtouch.step(state, actions, turn=800)
|
||||||
|
|
||||||
|
|
||||||
|
class PolicyAgent:
|
||||||
|
"""Direct policy agent that samples actions as in training/evaluation.
|
||||||
|
|
||||||
|
On each turn it computes the masked policy logits over all legal moves and
|
||||||
|
samples one action with `jrandom.categorical`, exactly mirroring
|
||||||
|
`PolicyValueNetwork.__call__` used during rollout collection.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
network: PolicyValueNetwork,
|
||||||
|
seed: int = 0,
|
||||||
|
):
|
||||||
|
self.network = network
|
||||||
|
self.key = jrandom.PRNGKey(seed)
|
||||||
|
self._infer = eqx.filter_jit(
|
||||||
|
lambda network, tensor, mask: network.policy_value(tensor, mask)
|
||||||
|
)
|
||||||
|
|
||||||
|
def warmup(self, height: int, width: int) -> None:
|
||||||
|
"""Compile shape-dependent network and competition transition kernels."""
|
||||||
|
grid = jnp.zeros((height, width), dtype=jnp.int32)
|
||||||
|
grid = grid.at[0, 0].set(1).at[height - 1, width - 1].set(2)
|
||||||
|
state = game.create_initial_state(grid)
|
||||||
|
state = state._replace(
|
||||||
|
armies=state.armies.at[0, 0].set(100).at[height - 1, width - 1].set(100),
|
||||||
|
time=jnp.int32(100),
|
||||||
|
)
|
||||||
|
obs = game.get_observation(state, 0)
|
||||||
|
self.policy_value(obs)
|
||||||
|
actions = jnp.stack([jnp.asarray(PASS_ACTION), jnp.asarray(PASS_ACTION)])
|
||||||
|
warmed_state, _ = competition_step(state, actions)
|
||||||
|
jax.block_until_ready(warmed_state)
|
||||||
|
|
||||||
|
def policy_value(self, obs: Observation):
|
||||||
|
mask = competition_move_mask(obs)
|
||||||
|
logits, value = self._infer(self.network, obs_to_tensor(obs), mask)
|
||||||
|
return np.asarray(logits), float(value), mask
|
||||||
|
|
||||||
|
def select(self, obs: Observation) -> tuple[np.ndarray, dict[str, float]]:
|
||||||
|
"""Sample one action from the masked policy and report statistics."""
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
logits, value, _ = self.policy_value(obs)
|
||||||
|
self.key, sample_key = jrandom.split(self.key)
|
||||||
|
index = int(jrandom.categorical(sample_key, jnp.asarray(logits)))
|
||||||
|
height, width = obs.armies.shape
|
||||||
|
action = _decode_policy_index(index, height, width)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
return action, {
|
||||||
|
"value": value,
|
||||||
|
"elapsed_ms": elapsed_ms,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Stdio driver
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _read_grid(stream, height: int):
|
||||||
|
return [[int(value) for value in stream.readline().split()] for _ in range(height)]
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args():
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
|
parser.add_argument("checkpoint", help="Equinox policy checkpoint")
|
||||||
|
parser.add_argument("--seed", type=int, default=0)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
args = parse_args()
|
||||||
|
handshake = sys.stdin.readline()
|
||||||
|
if not handshake:
|
||||||
|
return
|
||||||
|
player_index, height, width = (int(value) for value in handshake.split())
|
||||||
|
|
||||||
|
network = PolicyValueNetwork(jrandom.PRNGKey(args.seed))
|
||||||
|
network = eqx.tree_deserialise_leaves(args.checkpoint, network)
|
||||||
|
agent = PolicyAgent(network, seed=args.seed)
|
||||||
|
agent.warmup(height, width)
|
||||||
|
print(f"[competition] warmup complete for {height}x{width}", file=sys.stderr, flush=True)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
scalar_line = sys.stdin.readline()
|
||||||
|
if not scalar_line:
|
||||||
|
return
|
||||||
|
timestep, own_land, own_army, opponent_land, opponent_army = (
|
||||||
|
int(value) for value in scalar_line.split()
|
||||||
|
)
|
||||||
|
type_grid = _read_grid(sys.stdin, height)
|
||||||
|
owner_grid = _read_grid(sys.stdin, height)
|
||||||
|
army_grid = _read_grid(sys.stdin, height)
|
||||||
|
observation = observation_from_wire(
|
||||||
|
timestep,
|
||||||
|
own_land,
|
||||||
|
own_army,
|
||||||
|
opponent_land,
|
||||||
|
opponent_army,
|
||||||
|
type_grid,
|
||||||
|
owner_grid,
|
||||||
|
army_grid,
|
||||||
|
)
|
||||||
|
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
action, stats = agent.select(observation)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
print(
|
||||||
|
f"[competition] turn={timestep} value={stats['value']:+.3f} "
|
||||||
|
f"elapsed_ms={elapsed_ms:.1f}",
|
||||||
|
file=sys.stderr,
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
print(" ".join(str(int(value)) for value in action), flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
BIN
competition/model.eqx
Normal file
BIN
competition/model.eqx
Normal file
Binary file not shown.
3
competition/requirements-vendor.txt
Normal file
3
competition/requirements-vendor.txt
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
equinox==0.13.8
|
||||||
|
jaxtyping==0.3.11
|
||||||
|
wadler-lindig==0.1.7
|
||||||
8
competition/run.sh
Executable file
8
competition/run.sh
Executable file
|
|
@ -0,0 +1,8 @@
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
# Competition entrypoint. The sandbox runs `run.sh` from the submission root;
|
||||||
|
# cd to our own directory so relative paths work regardless of the cwd.
|
||||||
|
# `python` is the sandbox's CPython 3.12 with jax/numpy pre-installed; equinox
|
||||||
|
# and its deps are installed by build.sh from wheelhouse/.
|
||||||
|
cd "$(dirname "$0")"
|
||||||
|
exec python -u competition.py model.eqx
|
||||||
BIN
competition/wheelhouse/equinox-0.13.8-py3-none-any.whl
Normal file
BIN
competition/wheelhouse/equinox-0.13.8-py3-none-any.whl
Normal file
Binary file not shown.
BIN
competition/wheelhouse/jaxtyping-0.3.11-py3-none-any.whl
Normal file
BIN
competition/wheelhouse/jaxtyping-0.3.11-py3-none-any.whl
Normal file
Binary file not shown.
BIN
competition/wheelhouse/wadler_lindig-0.1.7-py3-none-any.whl
Normal file
BIN
competition/wheelhouse/wadler_lindig-0.1.7-py3-none-any.whl
Normal file
Binary file not shown.
BIN
ppo_model.eqx
BIN
ppo_model.eqx
Binary file not shown.
BIN
ppo_model.eqx.7
Normal file
BIN
ppo_model.eqx.7
Normal file
Binary file not shown.
Binary file not shown.
BIN
ppo_model_h100.eqx.2
Normal file
BIN
ppo_model_h100.eqx.2
Normal file
Binary file not shown.
BIN
ppo_model_h100.eqx.3
Normal file
BIN
ppo_model_h100.eqx.3
Normal file
Binary file not shown.
BIN
ppo_model_h100.eqx.4
Normal file
BIN
ppo_model_h100.eqx.4
Normal file
Binary file not shown.
11
requirements-competition.txt
Normal file
11
requirements-competition.txt
Normal file
|
|
@ -0,0 +1,11 @@
|
||||||
|
numpy==2.4.6
|
||||||
|
scipy==1.18.0
|
||||||
|
pandas==3.0.5
|
||||||
|
scikit-learn==1.9.0
|
||||||
|
# On Linux, the default PyPI wheel bundles CUDA. For a CPU-only build, uncomment:
|
||||||
|
jax==0.11.0
|
||||||
|
numba==0.66.0
|
||||||
|
networkx==3.6.1
|
||||||
|
safetensors==0.8.0
|
||||||
|
gymnasium==1.3.0
|
||||||
|
# torch==2.13.0
|
||||||
3
requirements-vendor.txt
Normal file
3
requirements-vendor.txt
Normal file
|
|
@ -0,0 +1,3 @@
|
||||||
|
equinox==0.13.8
|
||||||
|
jaxtyping==0.3.11
|
||||||
|
wadler-lindig==0.1.7
|
||||||
342
scripts/competition.py
Normal file
342
scripts/competition.py
Normal file
|
|
@ -0,0 +1,342 @@
|
||||||
|
"""Competition stdio agent: direct policy sampling over a trained equinox checkpoint.
|
||||||
|
|
||||||
|
Self-contained competition bot. It speaks the wire protocol in
|
||||||
|
`generals-bots/competition/protocol.py` (handshake, then one observation frame
|
||||||
|
per turn, one action line per turn, EOF on stdin = game over) and, on each
|
||||||
|
turn, samples an action from the policy network exactly as the model is used
|
||||||
|
during training/evaluation — no tree search.
|
||||||
|
|
||||||
|
Unlike `scripts/mcts_agent.py`, this module does **not** import the
|
||||||
|
`general_bots_training` training package. It inlines the policy-value network
|
||||||
|
and the observation encoding so the script can be dropped into the competition
|
||||||
|
sandbox with just the engine (`generals`), `equinox`, `jax`, and `numpy`
|
||||||
|
available.
|
||||||
|
|
||||||
|
This is the search-free counterpart to `scripts/competition_puct.py`: it shares
|
||||||
|
the same network and observation encoding, but selects an action by sampling
|
||||||
|
from the masked policy (as in `PolicyValueNetwork.__call__`) instead of running
|
||||||
|
particle PUCT.
|
||||||
|
|
||||||
|
Run it directly through the bundled matchup driver, e.g.:
|
||||||
|
|
||||||
|
PYTHONPATH=src:generals-bots .venv/bin/python \
|
||||||
|
generals-bots/competition/matchup.py \
|
||||||
|
scripts/competition.py --checkpoint ppo_model.eqx \
|
||||||
|
generals-bots/competition/agents/expander_python/run.sh \
|
||||||
|
--mode competition
|
||||||
|
|
||||||
|
or wrap it in a `run.sh` that passes the checkpoint path.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
|
||||||
|
import equinox as eqx
|
||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import jax.random as jrandom
|
||||||
|
import numpy as np
|
||||||
|
from generals.core import game
|
||||||
|
from generals.core.action import compute_valid_move_mask
|
||||||
|
from generals.core.game import GameState
|
||||||
|
from generals.core.observation import Observation
|
||||||
|
from generals.modifiers import build_castles, deathtouch
|
||||||
|
|
||||||
|
os.environ.setdefault("JAX_PLATFORMS", "cpu")
|
||||||
|
|
||||||
|
PASS_ACTION = np.array([1, 0, 0, 0, 0], dtype=np.int32)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Policy-value network (inlined from general_bots_training.network)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def obs_to_tensor(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Encode an Observation into a (C, H, W) float32 tensor for the network.
|
||||||
|
|
||||||
|
Armies and army-counts are log-normalized; land counts are divided by the
|
||||||
|
number of cells; scalar values are broadcast to spatial planes so a plain
|
||||||
|
conv stack can consume them. Grid-size agnostic.
|
||||||
|
"""
|
||||||
|
H, W = obs.armies.shape
|
||||||
|
armies = jnp.log1p(obs.armies.astype(jnp.float32)) / jnp.log(50.0)
|
||||||
|
|
||||||
|
def bcast(scalar):
|
||||||
|
return jnp.broadcast_to(scalar.astype(jnp.float32), (H, W))
|
||||||
|
|
||||||
|
own_land = bcast(obs.owned_land_count / (H * W))
|
||||||
|
own_army = bcast(jnp.log1p(obs.owned_army_count.astype(jnp.float32)) / jnp.log(50.0))
|
||||||
|
opp_land = bcast(obs.opponent_land_count / (H * W))
|
||||||
|
opp_army = bcast(jnp.log1p(obs.opponent_army_count.astype(jnp.float32)) / jnp.log(50.0))
|
||||||
|
timestep = bcast(jnp.log1p(obs.timestep.astype(jnp.float32)) / jnp.log(1201.0))
|
||||||
|
|
||||||
|
return jnp.stack(
|
||||||
|
[
|
||||||
|
armies,
|
||||||
|
obs.generals.astype(jnp.float32),
|
||||||
|
obs.castles.astype(jnp.float32),
|
||||||
|
obs.mountains.astype(jnp.float32),
|
||||||
|
obs.neutral_cells.astype(jnp.float32),
|
||||||
|
obs.owned_cells.astype(jnp.float32),
|
||||||
|
obs.opponent_cells.astype(jnp.float32),
|
||||||
|
obs.fog_cells.astype(jnp.float32),
|
||||||
|
obs.structures_in_fog.astype(jnp.float32),
|
||||||
|
own_land,
|
||||||
|
own_army,
|
||||||
|
opp_land,
|
||||||
|
opp_army,
|
||||||
|
timestep,
|
||||||
|
],
|
||||||
|
axis=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PolicyValueNetwork(eqx.Module):
|
||||||
|
"""Conv policy-value network.
|
||||||
|
|
||||||
|
Action layout: 4 full-move directions and 4 half-move (split) directions
|
||||||
|
per source cell, plus one global pass action. Invalid moves are masked to
|
||||||
|
-1e9; pass is always available.
|
||||||
|
|
||||||
|
The value head uses global average pooling over the spatial grid, so the
|
||||||
|
network works for any board size without reshaping linear layers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
conv1: eqx.nn.Conv2d
|
||||||
|
conv2: eqx.nn.Conv2d
|
||||||
|
conv3: eqx.nn.Conv2d
|
||||||
|
conv4: eqx.nn.Conv2d
|
||||||
|
policy_conv: eqx.nn.Conv2d
|
||||||
|
value_conv: eqx.nn.Conv2d
|
||||||
|
value_linear1: eqx.nn.Linear
|
||||||
|
value_linear2: eqx.nn.Linear
|
||||||
|
|
||||||
|
def __init__(self, key, in_channels: int = 14, channels=(32, 32, 32, 16)):
|
||||||
|
keys = jrandom.split(key, 8)
|
||||||
|
self.conv1 = eqx.nn.Conv2d(in_channels, channels[0], kernel_size=3, padding=1, key=keys[0])
|
||||||
|
self.conv2 = eqx.nn.Conv2d(channels[0], channels[1], kernel_size=3, padding=1, key=keys[1])
|
||||||
|
self.conv3 = eqx.nn.Conv2d(channels[1], channels[2], kernel_size=3, padding=1, key=keys[2])
|
||||||
|
self.conv4 = eqx.nn.Conv2d(channels[2], channels[3], kernel_size=3, padding=1, key=keys[3])
|
||||||
|
# 9 = 4 dirs (full) + 4 dirs (half) + 1 pass
|
||||||
|
self.policy_conv = eqx.nn.Conv2d(channels[3], 9, kernel_size=1, key=keys[4])
|
||||||
|
self.value_conv = eqx.nn.Conv2d(channels[3], 4, kernel_size=1, key=keys[5])
|
||||||
|
self.value_linear1 = eqx.nn.Linear(4, 64, key=keys[6])
|
||||||
|
self.value_linear2 = eqx.nn.Linear(64, 1, key=keys[7])
|
||||||
|
|
||||||
|
def _features(self, obs):
|
||||||
|
x = jax.nn.relu(self.conv1(obs))
|
||||||
|
x = jax.nn.relu(self.conv2(x))
|
||||||
|
x = jax.nn.relu(self.conv3(x))
|
||||||
|
x = jax.nn.relu(self.conv4(x))
|
||||||
|
return x
|
||||||
|
|
||||||
|
def _value_from_features(self, feat):
|
||||||
|
v = jax.nn.relu(self.value_conv(feat)) # (4, H, W)
|
||||||
|
v = v.mean(axis=(1, 2)) # (4,) global average pool
|
||||||
|
v = jax.nn.relu(self.value_linear1(v))
|
||||||
|
return self.value_linear2(v)[0]
|
||||||
|
|
||||||
|
def policy_value(self, obs, mask):
|
||||||
|
"""Return deterministic masked policy logits and the critic value."""
|
||||||
|
features = self._features(obs)
|
||||||
|
return self._policy_logits(features, mask), self._value_from_features(features)
|
||||||
|
|
||||||
|
def _policy_logits(self, feat, mask):
|
||||||
|
"""Return 8*H*W move logits plus one global pass logit."""
|
||||||
|
logits = self.policy_conv(feat) # (9, H, W)
|
||||||
|
mask_t = jnp.transpose(mask, (2, 0, 1)) # (4, H, W)
|
||||||
|
penalty = (1.0 - mask_t) * -1e9
|
||||||
|
move_penalty = jnp.concatenate([penalty, penalty], axis=0)
|
||||||
|
move_logits = (logits[:8] + move_penalty).reshape(-1)
|
||||||
|
pass_logit = jnp.mean(logits[8])[None]
|
||||||
|
return jnp.concatenate([move_logits, pass_logit])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Action selection (matches training/evaluation rollout)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def competition_move_mask(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Legal move mask, identical to the one used during training/evaluation."""
|
||||||
|
return compute_valid_move_mask(obs.armies, obs.owned_cells, obs.mountains)
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_policy_index(index: int, height: int, width: int) -> np.ndarray:
|
||||||
|
cells = height * width
|
||||||
|
if index == 8 * cells:
|
||||||
|
return PASS_ACTION.copy()
|
||||||
|
encoded_direction, position = divmod(index, cells)
|
||||||
|
row, col = divmod(position, width)
|
||||||
|
split = int(encoded_direction >= 4)
|
||||||
|
direction = encoded_direction - 4 if split else encoded_direction
|
||||||
|
return np.array([0, row, col, direction, split], dtype=np.int32)
|
||||||
|
|
||||||
|
|
||||||
|
def observation_from_wire(
|
||||||
|
timestep: int,
|
||||||
|
owned_land_count: int,
|
||||||
|
owned_army_count: int,
|
||||||
|
opponent_land_count: int,
|
||||||
|
opponent_army_count: int,
|
||||||
|
type_grid,
|
||||||
|
owner_grid,
|
||||||
|
army_grid,
|
||||||
|
) -> Observation:
|
||||||
|
"""Convert a competition wire frame into the network's Observation type."""
|
||||||
|
types = jnp.asarray(type_grid, dtype=jnp.int32)
|
||||||
|
owners = jnp.asarray(owner_grid, dtype=jnp.int32)
|
||||||
|
armies = jnp.asarray(army_grid, dtype=jnp.int32)
|
||||||
|
visible = (types != 0) & (types != 5)
|
||||||
|
return Observation(
|
||||||
|
armies=armies,
|
||||||
|
generals=types == 4,
|
||||||
|
castles=types == 3,
|
||||||
|
mountains=types == 2,
|
||||||
|
neutral_cells=visible & (owners == 0) & (types != 2),
|
||||||
|
owned_cells=owners == 1,
|
||||||
|
opponent_cells=owners == 2,
|
||||||
|
fog_cells=types == 0,
|
||||||
|
structures_in_fog=types == 5,
|
||||||
|
owned_land_count=jnp.int32(owned_land_count),
|
||||||
|
owned_army_count=jnp.int32(owned_army_count),
|
||||||
|
opponent_land_count=jnp.int32(opponent_land_count),
|
||||||
|
opponent_army_count=jnp.int32(opponent_army_count),
|
||||||
|
timestep=jnp.int32(timestep),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def competition_step(state: GameState, actions: jnp.ndarray):
|
||||||
|
state, actions = build_castles.apply_build_actions(state, actions)
|
||||||
|
return deathtouch.step(state, actions, turn=800)
|
||||||
|
|
||||||
|
|
||||||
|
class PolicyAgent:
|
||||||
|
"""Direct policy agent that samples actions as in training/evaluation.
|
||||||
|
|
||||||
|
On each turn it computes the masked policy logits over all legal moves and
|
||||||
|
samples one action with `jrandom.categorical`, exactly mirroring
|
||||||
|
`PolicyValueNetwork.__call__` used during rollout collection.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
network: PolicyValueNetwork,
|
||||||
|
seed: int = 0,
|
||||||
|
):
|
||||||
|
self.network = network
|
||||||
|
self.key = jrandom.PRNGKey(seed)
|
||||||
|
self._infer = eqx.filter_jit(
|
||||||
|
lambda network, tensor, mask: network.policy_value(tensor, mask)
|
||||||
|
)
|
||||||
|
|
||||||
|
def warmup(self, height: int, width: int) -> None:
|
||||||
|
"""Compile shape-dependent network and competition transition kernels."""
|
||||||
|
grid = jnp.zeros((height, width), dtype=jnp.int32)
|
||||||
|
grid = grid.at[0, 0].set(1).at[height - 1, width - 1].set(2)
|
||||||
|
state = game.create_initial_state(grid)
|
||||||
|
state = state._replace(
|
||||||
|
armies=state.armies.at[0, 0].set(100).at[height - 1, width - 1].set(100),
|
||||||
|
time=jnp.int32(100),
|
||||||
|
)
|
||||||
|
obs = game.get_observation(state, 0)
|
||||||
|
self.policy_value(obs)
|
||||||
|
actions = jnp.stack([jnp.asarray(PASS_ACTION), jnp.asarray(PASS_ACTION)])
|
||||||
|
warmed_state, _ = competition_step(state, actions)
|
||||||
|
jax.block_until_ready(warmed_state)
|
||||||
|
|
||||||
|
def policy_value(self, obs: Observation):
|
||||||
|
mask = competition_move_mask(obs)
|
||||||
|
logits, value = self._infer(self.network, obs_to_tensor(obs), mask)
|
||||||
|
return np.asarray(logits), float(value), mask
|
||||||
|
|
||||||
|
def select(self, obs: Observation) -> tuple[np.ndarray, dict[str, float]]:
|
||||||
|
"""Sample one action from the masked policy and report statistics."""
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
logits, value, _ = self.policy_value(obs)
|
||||||
|
self.key, sample_key = jrandom.split(self.key)
|
||||||
|
index = int(jrandom.categorical(sample_key, jnp.asarray(logits)))
|
||||||
|
height, width = obs.armies.shape
|
||||||
|
action = _decode_policy_index(index, height, width)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
return action, {
|
||||||
|
"value": value,
|
||||||
|
"elapsed_ms": elapsed_ms,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Stdio driver
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _read_grid(stream, height: int):
|
||||||
|
return [[int(value) for value in stream.readline().split()] for _ in range(height)]
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args():
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
|
parser.add_argument("checkpoint", help="Equinox policy checkpoint")
|
||||||
|
# Accepted for compatibility with the original run.sh; the direct policy
|
||||||
|
# agent does no search, so these are ignored.
|
||||||
|
parser.add_argument("--time-budget-ms", type=float, default=125.0)
|
||||||
|
parser.add_argument("--max-simulations", type=int, default=128)
|
||||||
|
parser.add_argument("--rollout-depth", type=int, default=2)
|
||||||
|
parser.add_argument("--top-k", type=int, default=20)
|
||||||
|
parser.add_argument("--seed", type=int, default=0)
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
args = parse_args()
|
||||||
|
handshake = sys.stdin.readline()
|
||||||
|
if not handshake:
|
||||||
|
return
|
||||||
|
player_index, height, width = (int(value) for value in handshake.split())
|
||||||
|
|
||||||
|
network = PolicyValueNetwork(jrandom.PRNGKey(args.seed))
|
||||||
|
network = eqx.tree_deserialise_leaves(args.checkpoint, network)
|
||||||
|
agent = PolicyAgent(network, seed=args.seed)
|
||||||
|
agent.warmup(height, width)
|
||||||
|
print(f"[competition] warmup complete for {height}x{width}", file=sys.stderr, flush=True)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
scalar_line = sys.stdin.readline()
|
||||||
|
if not scalar_line:
|
||||||
|
return
|
||||||
|
timestep, own_land, own_army, opponent_land, opponent_army = (
|
||||||
|
int(value) for value in scalar_line.split()
|
||||||
|
)
|
||||||
|
type_grid = _read_grid(sys.stdin, height)
|
||||||
|
owner_grid = _read_grid(sys.stdin, height)
|
||||||
|
army_grid = _read_grid(sys.stdin, height)
|
||||||
|
observation = observation_from_wire(
|
||||||
|
timestep,
|
||||||
|
own_land,
|
||||||
|
own_army,
|
||||||
|
opponent_land,
|
||||||
|
opponent_army,
|
||||||
|
type_grid,
|
||||||
|
owner_grid,
|
||||||
|
army_grid,
|
||||||
|
)
|
||||||
|
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
action, stats = agent.select(observation)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
print(
|
||||||
|
f"[competition] turn={timestep} value={stats['value']:+.3f} "
|
||||||
|
f"elapsed_ms={elapsed_ms:.1f}",
|
||||||
|
file=sys.stderr,
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
print(" ".join(str(int(value)) for value in action), flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
552
scripts/competition_puct.py
Normal file
552
scripts/competition_puct.py
Normal file
|
|
@ -0,0 +1,552 @@
|
||||||
|
"""Competition stdio agent: particle PUCT over a trained equinox checkpoint.
|
||||||
|
|
||||||
|
Self-contained competition bot. It speaks the wire protocol in
|
||||||
|
`generals-bots/competition/protocol.py` (handshake, then one observation frame
|
||||||
|
per turn, one action line per turn, EOF on stdin = game over) and runs
|
||||||
|
deadline-bounded root PUCT using only the perspective-relative wire
|
||||||
|
observation.
|
||||||
|
|
||||||
|
Unlike `scripts/mcts_agent.py`, this module does **not** import the
|
||||||
|
`general_bots_training` training package. It inlines the policy-value network,
|
||||||
|
the observation encoding, and the particle search so the script can be dropped
|
||||||
|
into the competition sandbox with just the engine (`generals`), `equinox`,
|
||||||
|
`jax`, and `numpy` available.
|
||||||
|
|
||||||
|
Run it directly through the bundled matchup driver, e.g.:
|
||||||
|
|
||||||
|
PYTHONPATH=src:generals-bots .venv/bin/python \
|
||||||
|
generals-bots/competition/matchup.py \
|
||||||
|
scripts/competition.py --checkpoint ppo_model.eqx \
|
||||||
|
generals-bots/competition/agents/expander_python/run.sh \
|
||||||
|
--mode competition
|
||||||
|
|
||||||
|
or wrap it in a `run.sh` that passes the checkpoint path.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import time
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import equinox as eqx
|
||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import jax.random as jrandom
|
||||||
|
import numpy as np
|
||||||
|
from generals.core import game
|
||||||
|
from generals.core.action import compute_valid_move_mask
|
||||||
|
from generals.core.game import GameState
|
||||||
|
from generals.core.observation import Observation
|
||||||
|
from generals.modifiers import build_castles, deathtouch
|
||||||
|
|
||||||
|
os.environ.setdefault("JAX_PLATFORMS", "cpu")
|
||||||
|
|
||||||
|
PASS_ACTION = np.array([1, 0, 0, 0, 0], dtype=np.int32)
|
||||||
|
|
||||||
|
# Search/agent defaults, previously exposed as CLI flags. Kept as constants so
|
||||||
|
# the script can be launched by a bare `run.sh` with no extra arguments.
|
||||||
|
TIME_BUDGET_MS = 125.0
|
||||||
|
MAX_SIMULATIONS = 128
|
||||||
|
ROLLOUT_DEPTH = 2
|
||||||
|
TOP_K = 20
|
||||||
|
SEED = 0
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Policy-value network (inlined from general_bots_training.network)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def obs_to_tensor(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Encode an Observation into a (C, H, W) float32 tensor for the network.
|
||||||
|
|
||||||
|
Armies and army-counts are log-normalized; land counts are divided by the
|
||||||
|
number of cells; scalar values are broadcast to spatial planes so a plain
|
||||||
|
conv stack can consume them. Grid-size agnostic.
|
||||||
|
"""
|
||||||
|
H, W = obs.armies.shape
|
||||||
|
armies = jnp.log1p(obs.armies.astype(jnp.float32)) / jnp.log(50.0)
|
||||||
|
|
||||||
|
def bcast(scalar):
|
||||||
|
return jnp.broadcast_to(scalar.astype(jnp.float32), (H, W))
|
||||||
|
|
||||||
|
own_land = bcast(obs.owned_land_count / (H * W))
|
||||||
|
own_army = bcast(jnp.log1p(obs.owned_army_count.astype(jnp.float32)) / jnp.log(50.0))
|
||||||
|
opp_land = bcast(obs.opponent_land_count / (H * W))
|
||||||
|
opp_army = bcast(jnp.log1p(obs.opponent_army_count.astype(jnp.float32)) / jnp.log(50.0))
|
||||||
|
timestep = bcast(jnp.log1p(obs.timestep.astype(jnp.float32)) / jnp.log(1201.0))
|
||||||
|
|
||||||
|
return jnp.stack(
|
||||||
|
[
|
||||||
|
armies,
|
||||||
|
obs.generals.astype(jnp.float32),
|
||||||
|
obs.castles.astype(jnp.float32),
|
||||||
|
obs.mountains.astype(jnp.float32),
|
||||||
|
obs.neutral_cells.astype(jnp.float32),
|
||||||
|
obs.owned_cells.astype(jnp.float32),
|
||||||
|
obs.opponent_cells.astype(jnp.float32),
|
||||||
|
obs.fog_cells.astype(jnp.float32),
|
||||||
|
obs.structures_in_fog.astype(jnp.float32),
|
||||||
|
own_land,
|
||||||
|
own_army,
|
||||||
|
opp_land,
|
||||||
|
opp_army,
|
||||||
|
timestep,
|
||||||
|
],
|
||||||
|
axis=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class PolicyValueNetwork(eqx.Module):
|
||||||
|
"""Conv policy-value network.
|
||||||
|
|
||||||
|
Action layout: 4 full-move directions and 4 half-move (split) directions
|
||||||
|
per source cell, plus one global pass action. Invalid moves are masked to
|
||||||
|
-1e9; pass is always available.
|
||||||
|
|
||||||
|
The value head uses global average pooling over the spatial grid, so the
|
||||||
|
network works for any board size without reshaping linear layers.
|
||||||
|
"""
|
||||||
|
|
||||||
|
conv1: eqx.nn.Conv2d
|
||||||
|
conv2: eqx.nn.Conv2d
|
||||||
|
conv3: eqx.nn.Conv2d
|
||||||
|
conv4: eqx.nn.Conv2d
|
||||||
|
policy_conv: eqx.nn.Conv2d
|
||||||
|
value_conv: eqx.nn.Conv2d
|
||||||
|
value_linear1: eqx.nn.Linear
|
||||||
|
value_linear2: eqx.nn.Linear
|
||||||
|
|
||||||
|
def __init__(self, key, in_channels: int = 14, channels=(32, 32, 32, 16)):
|
||||||
|
keys = jrandom.split(key, 8)
|
||||||
|
self.conv1 = eqx.nn.Conv2d(in_channels, channels[0], kernel_size=3, padding=1, key=keys[0])
|
||||||
|
self.conv2 = eqx.nn.Conv2d(channels[0], channels[1], kernel_size=3, padding=1, key=keys[1])
|
||||||
|
self.conv3 = eqx.nn.Conv2d(channels[1], channels[2], kernel_size=3, padding=1, key=keys[2])
|
||||||
|
self.conv4 = eqx.nn.Conv2d(channels[2], channels[3], kernel_size=3, padding=1, key=keys[3])
|
||||||
|
# 9 = 4 dirs (full) + 4 dirs (half) + 1 pass
|
||||||
|
self.policy_conv = eqx.nn.Conv2d(channels[3], 9, kernel_size=1, key=keys[4])
|
||||||
|
self.value_conv = eqx.nn.Conv2d(channels[3], 4, kernel_size=1, key=keys[5])
|
||||||
|
self.value_linear1 = eqx.nn.Linear(4, 64, key=keys[6])
|
||||||
|
self.value_linear2 = eqx.nn.Linear(64, 1, key=keys[7])
|
||||||
|
|
||||||
|
def _features(self, obs):
|
||||||
|
x = jax.nn.relu(self.conv1(obs))
|
||||||
|
x = jax.nn.relu(self.conv2(x))
|
||||||
|
x = jax.nn.relu(self.conv3(x))
|
||||||
|
x = jax.nn.relu(self.conv4(x))
|
||||||
|
return x
|
||||||
|
|
||||||
|
def _value_from_features(self, feat):
|
||||||
|
v = jax.nn.relu(self.value_conv(feat)) # (4, H, W)
|
||||||
|
v = v.mean(axis=(1, 2)) # (4,) global average pool
|
||||||
|
v = jax.nn.relu(self.value_linear1(v))
|
||||||
|
return self.value_linear2(v)[0]
|
||||||
|
|
||||||
|
def policy_value(self, obs, mask):
|
||||||
|
"""Return deterministic masked policy logits and the critic value."""
|
||||||
|
features = self._features(obs)
|
||||||
|
return self._policy_logits(features, mask), self._value_from_features(features)
|
||||||
|
|
||||||
|
def _policy_logits(self, feat, mask):
|
||||||
|
"""Return 8*H*W move logits plus one global pass logit."""
|
||||||
|
logits = self.policy_conv(feat) # (9, H, W)
|
||||||
|
mask_t = jnp.transpose(mask, (2, 0, 1)) # (4, H, W)
|
||||||
|
penalty = (1.0 - mask_t) * -1e9
|
||||||
|
move_penalty = jnp.concatenate([penalty, penalty], axis=0)
|
||||||
|
move_logits = (logits[:8] + move_penalty).reshape(-1)
|
||||||
|
pass_logit = jnp.mean(logits[8])[None]
|
||||||
|
return jnp.concatenate([move_logits, pass_logit])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Particle PUCT search (inlined from general_bots_training.mcts)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class MCTSConfig:
|
||||||
|
time_budget_ms: float = 100.0
|
||||||
|
max_simulations: int = 128
|
||||||
|
rollout_depth: int = 2
|
||||||
|
top_k: int = 20
|
||||||
|
opponent_top_k: int = 12
|
||||||
|
max_build_actions: int = 3
|
||||||
|
c_puct: float = 1.5
|
||||||
|
value_scale: float = 5.0
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def competition_move_mask(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Legal move mask that conservatively blocks unresolved fog structures."""
|
||||||
|
blocked = obs.mountains | obs.structures_in_fog
|
||||||
|
return compute_valid_move_mask(obs.armies, obs.owned_cells, blocked)
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def build_cost_grid(obs: Observation) -> jnp.ndarray:
|
||||||
|
"""Exact own-castle build costs derivable from a fog observation."""
|
||||||
|
structures = ((obs.castles | obs.generals) & obs.owned_cells).astype(jnp.int32)
|
||||||
|
height, width = structures.shape
|
||||||
|
radius = 6
|
||||||
|
padded = jnp.pad(structures, radius)
|
||||||
|
costs = jnp.full((height, width), 35, dtype=jnp.int32)
|
||||||
|
for row_offset in range(-radius, radius + 1):
|
||||||
|
for col_offset in range(-radius, radius + 1):
|
||||||
|
surcharge = 14 - 2 * (abs(row_offset) + abs(col_offset))
|
||||||
|
if surcharge > 0:
|
||||||
|
shifted = padded[
|
||||||
|
radius + row_offset : radius + row_offset + height,
|
||||||
|
radius + col_offset : radius + col_offset + width,
|
||||||
|
]
|
||||||
|
costs = costs + surcharge * shifted
|
||||||
|
return costs
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def valid_build_mask(obs: Observation) -> jnp.ndarray:
|
||||||
|
costs = build_cost_grid(obs)
|
||||||
|
plain_owned = obs.owned_cells & ~obs.generals & ~obs.castles
|
||||||
|
return plain_owned & (obs.armies >= costs)
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_policy_index(index: int, height: int, width: int) -> np.ndarray:
|
||||||
|
cells = height * width
|
||||||
|
if index == 8 * cells:
|
||||||
|
return PASS_ACTION.copy()
|
||||||
|
encoded_direction, position = divmod(index, cells)
|
||||||
|
row, col = divmod(position, width)
|
||||||
|
split = int(encoded_direction >= 4)
|
||||||
|
direction = encoded_direction - 4 if split else encoded_direction
|
||||||
|
return np.array([0, row, col, direction, split], dtype=np.int32)
|
||||||
|
|
||||||
|
|
||||||
|
def observation_from_wire(
|
||||||
|
timestep: int,
|
||||||
|
owned_land_count: int,
|
||||||
|
owned_army_count: int,
|
||||||
|
opponent_land_count: int,
|
||||||
|
opponent_army_count: int,
|
||||||
|
type_grid,
|
||||||
|
owner_grid,
|
||||||
|
army_grid,
|
||||||
|
) -> Observation:
|
||||||
|
"""Convert a competition wire frame into the network's Observation type."""
|
||||||
|
types = jnp.asarray(type_grid, dtype=jnp.int32)
|
||||||
|
owners = jnp.asarray(owner_grid, dtype=jnp.int32)
|
||||||
|
armies = jnp.asarray(army_grid, dtype=jnp.int32)
|
||||||
|
visible = (types != 0) & (types != 5)
|
||||||
|
return Observation(
|
||||||
|
armies=armies,
|
||||||
|
generals=types == 4,
|
||||||
|
castles=types == 3,
|
||||||
|
mountains=types == 2,
|
||||||
|
neutral_cells=visible & (owners == 0) & (types != 2),
|
||||||
|
owned_cells=owners == 1,
|
||||||
|
opponent_cells=owners == 2,
|
||||||
|
fog_cells=types == 0,
|
||||||
|
structures_in_fog=types == 5,
|
||||||
|
owned_land_count=jnp.int32(owned_land_count),
|
||||||
|
owned_army_count=jnp.int32(owned_army_count),
|
||||||
|
opponent_land_count=jnp.int32(opponent_land_count),
|
||||||
|
opponent_army_count=jnp.int32(opponent_army_count),
|
||||||
|
timestep=jnp.int32(timestep),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def sample_determinization(
|
||||||
|
obs: Observation, player_index: int, rng: np.random.Generator
|
||||||
|
) -> GameState:
|
||||||
|
"""Sample a conservative hidden state consistent with visible cells/totals."""
|
||||||
|
armies = np.asarray(obs.armies, dtype=np.int32).copy()
|
||||||
|
owned = np.asarray(obs.owned_cells, dtype=bool)
|
||||||
|
visible_opponent = np.asarray(obs.opponent_cells, dtype=bool)
|
||||||
|
mountains = np.asarray(obs.mountains | obs.structures_in_fog, dtype=bool)
|
||||||
|
castles = np.asarray(obs.castles, dtype=bool)
|
||||||
|
generals = np.asarray(obs.generals, dtype=bool).copy()
|
||||||
|
fog = np.asarray(obs.fog_cells, dtype=bool) & ~mountains
|
||||||
|
|
||||||
|
opponent = visible_opponent.copy()
|
||||||
|
hidden_candidates = np.argwhere(fog & ~owned)
|
||||||
|
missing_land = max(0, int(obs.opponent_land_count) - int(opponent.sum()))
|
||||||
|
if len(hidden_candidates):
|
||||||
|
chosen = hidden_candidates[
|
||||||
|
rng.choice(
|
||||||
|
len(hidden_candidates),
|
||||||
|
size=min(missing_land, len(hidden_candidates)),
|
||||||
|
replace=False,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
opponent[chosen[:, 0], chosen[:, 1]] = True
|
||||||
|
|
||||||
|
visible_enemy_general = generals & opponent
|
||||||
|
if not visible_enemy_general.any():
|
||||||
|
candidates = np.argwhere((opponent | fog) & ~owned & ~mountains)
|
||||||
|
if len(candidates):
|
||||||
|
row, col = candidates[rng.integers(len(candidates))]
|
||||||
|
generals[row, col] = True
|
||||||
|
opponent[row, col] = True
|
||||||
|
|
||||||
|
hidden_opponent = opponent & ~visible_opponent
|
||||||
|
remaining_army = max(0, int(obs.opponent_army_count) - int(armies[visible_opponent].sum()))
|
||||||
|
hidden_cells = np.argwhere(hidden_opponent)
|
||||||
|
if len(hidden_cells):
|
||||||
|
base = min(remaining_army, len(hidden_cells))
|
||||||
|
armies[hidden_opponent] = 0
|
||||||
|
armies[hidden_cells[:base, 0], hidden_cells[:base, 1]] = 1
|
||||||
|
remaining_army -= base
|
||||||
|
if remaining_army:
|
||||||
|
allocations = rng.multinomial(
|
||||||
|
remaining_army, np.full(len(hidden_cells), 1 / len(hidden_cells))
|
||||||
|
)
|
||||||
|
armies[hidden_cells[:, 0], hidden_cells[:, 1]] += allocations.astype(np.int32)
|
||||||
|
|
||||||
|
ownership_relative = np.stack([owned, opponent])
|
||||||
|
if player_index == 1:
|
||||||
|
ownership = ownership_relative[::-1]
|
||||||
|
else:
|
||||||
|
ownership = ownership_relative
|
||||||
|
passable = ~mountains
|
||||||
|
neutral = passable & ~ownership[0] & ~ownership[1]
|
||||||
|
|
||||||
|
general_positions = []
|
||||||
|
for absolute_player in range(2):
|
||||||
|
positions = np.argwhere(generals & ownership[absolute_player])
|
||||||
|
if len(positions):
|
||||||
|
general_positions.append(positions[0])
|
||||||
|
else:
|
||||||
|
fallback = np.argwhere(ownership[absolute_player] & passable)
|
||||||
|
general_positions.append(fallback[0] if len(fallback) else np.array([0, 0]))
|
||||||
|
|
||||||
|
return GameState(
|
||||||
|
armies=jnp.asarray(armies),
|
||||||
|
ownership=jnp.asarray(ownership),
|
||||||
|
ownership_neutral=jnp.asarray(neutral),
|
||||||
|
generals=jnp.asarray(generals),
|
||||||
|
castles=jnp.asarray(castles),
|
||||||
|
mountains=jnp.asarray(mountains),
|
||||||
|
passable=jnp.asarray(passable),
|
||||||
|
general_positions=jnp.asarray(general_positions, dtype=jnp.int32),
|
||||||
|
time=jnp.asarray(obs.timestep, dtype=jnp.int32),
|
||||||
|
winner=jnp.int32(-1),
|
||||||
|
pool_idx=jnp.int32(0),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def competition_step(state: GameState, actions: jnp.ndarray):
|
||||||
|
state, actions = build_castles.apply_build_actions(state, actions)
|
||||||
|
return deathtouch.step(state, actions, turn=800)
|
||||||
|
|
||||||
|
|
||||||
|
class ParticlePUCT:
|
||||||
|
"""Deadline-bounded root PUCT with policy-guided simultaneous rollouts."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
network: PolicyValueNetwork,
|
||||||
|
player_index: int,
|
||||||
|
config: MCTSConfig = MCTSConfig(),
|
||||||
|
seed: int = 0,
|
||||||
|
):
|
||||||
|
self.network = network
|
||||||
|
self.player_index = player_index
|
||||||
|
self.config = config
|
||||||
|
self.rng = np.random.default_rng(seed)
|
||||||
|
self._infer = eqx.filter_jit(
|
||||||
|
lambda network, tensor, mask: network.policy_value(tensor, mask)
|
||||||
|
)
|
||||||
|
|
||||||
|
def warmup(self, height: int, width: int) -> None:
|
||||||
|
"""Compile shape-dependent network and competition transition kernels."""
|
||||||
|
grid = jnp.zeros((height, width), dtype=jnp.int32)
|
||||||
|
grid = grid.at[0, 0].set(1).at[height - 1, width - 1].set(2)
|
||||||
|
state = game.create_initial_state(grid)
|
||||||
|
state = state._replace(
|
||||||
|
armies=state.armies.at[0, 0].set(100).at[height - 1, width - 1].set(100),
|
||||||
|
time=jnp.int32(100),
|
||||||
|
)
|
||||||
|
obs = game.get_observation(state, self.player_index)
|
||||||
|
self.policy_value(obs)
|
||||||
|
actions = jnp.stack([jnp.asarray(PASS_ACTION), jnp.asarray(PASS_ACTION)])
|
||||||
|
warmed_state, _ = competition_step(state, actions)
|
||||||
|
jax.block_until_ready(warmed_state)
|
||||||
|
candidates, _ = self.candidates(obs)
|
||||||
|
self._simulate(obs, candidates[0])
|
||||||
|
self._simulate(obs, candidates[min(1, len(candidates) - 1)])
|
||||||
|
|
||||||
|
def policy_value(self, obs: Observation):
|
||||||
|
mask = competition_move_mask(obs)
|
||||||
|
logits, value = self._infer(self.network, obs_to_tensor(obs), mask)
|
||||||
|
return np.asarray(logits), float(value), mask
|
||||||
|
|
||||||
|
def candidates(self, obs: Observation, top_k: int | None = None):
|
||||||
|
logits, _, _ = self.policy_value(obs)
|
||||||
|
height, width = obs.armies.shape
|
||||||
|
valid_indices = np.flatnonzero(logits > -1e8)
|
||||||
|
count = min(top_k or self.config.top_k, len(valid_indices))
|
||||||
|
ranked = np.argsort(logits[valid_indices])[::-1][:count]
|
||||||
|
indices = valid_indices[ranked]
|
||||||
|
selected_logits = logits[indices]
|
||||||
|
selected_logits = selected_logits - selected_logits.max()
|
||||||
|
priors = np.exp(selected_logits)
|
||||||
|
actions = [_decode_policy_index(int(index), height, width) for index in indices]
|
||||||
|
|
||||||
|
build_mask = np.asarray(valid_build_mask(obs))
|
||||||
|
build_positions = np.argwhere(build_mask)
|
||||||
|
if len(build_positions):
|
||||||
|
costs = np.asarray(build_cost_grid(obs))
|
||||||
|
armies = np.asarray(obs.armies)
|
||||||
|
scores = np.array([armies[r, c] - costs[r, c] for r, c in build_positions])
|
||||||
|
order = np.argsort(scores)[::-1][: self.config.max_build_actions]
|
||||||
|
build_prior = max(float(priors.sum()) * 0.05, 1e-3)
|
||||||
|
for position_index in order:
|
||||||
|
row, col = build_positions[position_index]
|
||||||
|
actions.append(np.array([2, row, col, 0, 0], dtype=np.int32))
|
||||||
|
priors = np.append(priors, build_prior)
|
||||||
|
|
||||||
|
priors = priors / priors.sum()
|
||||||
|
return actions, priors
|
||||||
|
|
||||||
|
def _sample_policy_action(self, obs: Observation, top_k: int) -> np.ndarray:
|
||||||
|
actions, priors = self.candidates(obs, top_k)
|
||||||
|
return actions[int(self.rng.choice(len(actions), p=priors))]
|
||||||
|
|
||||||
|
def _leaf_value(self, state: GameState, info) -> float:
|
||||||
|
if bool(info.is_done):
|
||||||
|
winner = int(info.winner)
|
||||||
|
return 0.0 if winner < 0 else (1.0 if winner == self.player_index else -1.0)
|
||||||
|
obs = game.get_observation(state, self.player_index)
|
||||||
|
_, value, _ = self.policy_value(obs)
|
||||||
|
return float(np.tanh(value / self.config.value_scale))
|
||||||
|
|
||||||
|
def _simulate(self, root_obs: Observation, root_action: np.ndarray) -> float:
|
||||||
|
state = sample_determinization(root_obs, self.player_index, self.rng)
|
||||||
|
info = game.get_info(state)
|
||||||
|
our_action = root_action
|
||||||
|
for _ in range(self.config.rollout_depth):
|
||||||
|
opponent_index = 1 - self.player_index
|
||||||
|
opponent_obs = game.get_observation(state, opponent_index)
|
||||||
|
opponent_action = self._sample_policy_action(opponent_obs, self.config.opponent_top_k)
|
||||||
|
joint_actions = [None, None]
|
||||||
|
joint_actions[self.player_index] = our_action
|
||||||
|
joint_actions[opponent_index] = opponent_action
|
||||||
|
state, info = competition_step(state, jnp.asarray(joint_actions, dtype=jnp.int32))
|
||||||
|
jax.block_until_ready(state)
|
||||||
|
if bool(info.is_done):
|
||||||
|
break
|
||||||
|
our_obs = game.get_observation(state, self.player_index)
|
||||||
|
our_action = self._sample_policy_action(our_obs, self.config.top_k)
|
||||||
|
return self._leaf_value(state, info)
|
||||||
|
|
||||||
|
def search(self, obs: Observation) -> tuple[np.ndarray, dict[str, float]]:
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
deadline = started_at + self.config.time_budget_ms / 1000.0
|
||||||
|
actions, priors = self.candidates(obs)
|
||||||
|
visits = np.zeros(len(actions), dtype=np.int32)
|
||||||
|
value_sums = np.zeros(len(actions), dtype=np.float64)
|
||||||
|
simulations = 0
|
||||||
|
estimated_simulation_seconds = 0.0
|
||||||
|
|
||||||
|
while (
|
||||||
|
simulations < self.config.max_simulations
|
||||||
|
and time.perf_counter() + estimated_simulation_seconds < deadline
|
||||||
|
):
|
||||||
|
total_visits = max(1, int(visits.sum()))
|
||||||
|
q_values = np.divide(
|
||||||
|
value_sums,
|
||||||
|
visits,
|
||||||
|
out=np.zeros_like(value_sums),
|
||||||
|
where=visits > 0,
|
||||||
|
)
|
||||||
|
scores = q_values + self.config.c_puct * priors * np.sqrt(total_visits) / (1 + visits)
|
||||||
|
action_index = int(np.argmax(scores))
|
||||||
|
simulation_started_at = time.perf_counter()
|
||||||
|
value = self._simulate(obs, actions[action_index])
|
||||||
|
simulation_elapsed = time.perf_counter() - simulation_started_at
|
||||||
|
estimated_simulation_seconds = max(estimated_simulation_seconds, simulation_elapsed)
|
||||||
|
visits[action_index] += 1
|
||||||
|
value_sums[action_index] += value
|
||||||
|
simulations += 1
|
||||||
|
|
||||||
|
selected = int(np.argmax(visits)) if simulations else int(np.argmax(priors))
|
||||||
|
return actions[selected], {
|
||||||
|
"simulations": float(simulations),
|
||||||
|
"selected_visits": float(visits[selected]),
|
||||||
|
"elapsed_budget_ms": self.config.time_budget_ms,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Stdio driver
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _read_grid(stream, height: int):
|
||||||
|
return [[int(value) for value in stream.readline().split()] for _ in range(height)]
|
||||||
|
|
||||||
|
|
||||||
|
def parse_args():
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
|
parser.add_argument("checkpoint", help="Equinox policy checkpoint")
|
||||||
|
return parser.parse_args()
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
args = parse_args()
|
||||||
|
handshake = sys.stdin.readline()
|
||||||
|
if not handshake:
|
||||||
|
return
|
||||||
|
player_index, height, width = (int(value) for value in handshake.split())
|
||||||
|
|
||||||
|
network = PolicyValueNetwork(jrandom.PRNGKey(SEED))
|
||||||
|
network = eqx.tree_deserialise_leaves(args.checkpoint, network)
|
||||||
|
search = ParticlePUCT(
|
||||||
|
network,
|
||||||
|
player_index,
|
||||||
|
MCTSConfig(
|
||||||
|
time_budget_ms=TIME_BUDGET_MS,
|
||||||
|
max_simulations=MAX_SIMULATIONS,
|
||||||
|
rollout_depth=ROLLOUT_DEPTH,
|
||||||
|
top_k=TOP_K,
|
||||||
|
),
|
||||||
|
seed=SEED,
|
||||||
|
)
|
||||||
|
search.warmup(height, width)
|
||||||
|
print(f"[competition] warmup complete for {height}x{width}", file=sys.stderr, flush=True)
|
||||||
|
|
||||||
|
while True:
|
||||||
|
scalar_line = sys.stdin.readline()
|
||||||
|
if not scalar_line:
|
||||||
|
return
|
||||||
|
timestep, own_land, own_army, opponent_land, opponent_army = (
|
||||||
|
int(value) for value in scalar_line.split()
|
||||||
|
)
|
||||||
|
type_grid = _read_grid(sys.stdin, height)
|
||||||
|
owner_grid = _read_grid(sys.stdin, height)
|
||||||
|
army_grid = _read_grid(sys.stdin, height)
|
||||||
|
observation = observation_from_wire(
|
||||||
|
timestep,
|
||||||
|
own_land,
|
||||||
|
own_army,
|
||||||
|
opponent_land,
|
||||||
|
opponent_army,
|
||||||
|
type_grid,
|
||||||
|
owner_grid,
|
||||||
|
army_grid,
|
||||||
|
)
|
||||||
|
|
||||||
|
started_at = time.perf_counter()
|
||||||
|
action, stats = search.search(observation)
|
||||||
|
elapsed_ms = (time.perf_counter() - started_at) * 1000
|
||||||
|
print(
|
||||||
|
f"[competition] turn={timestep} simulations={int(stats['simulations'])} "
|
||||||
|
f"elapsed_ms={elapsed_ms:.1f}",
|
||||||
|
file=sys.stderr,
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
print(" ".join(str(int(value)) for value in action), flush=True)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
4
scripts/run.sh
Normal file
4
scripts/run.sh
Normal file
|
|
@ -0,0 +1,4 @@
|
||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
cd "$(dirname "$0")"
|
||||||
|
exec python -u competition.py ../ppo_model.eqx
|
||||||
Loading…
Reference in a new issue