Training a Pyramid Assisted Player by Imitation¶
Author: Rob Hendriks
Last verified: 17 September 2026
This tutorial bootstraps the trainable Pyramid (Pyr) assisted players from the known-correct classical Pyramid strategy.
As in the linear tutorial, we first train the learned layers by imitation. The difference is that the Pyramid strategy is recursive: every level reduces the problem before passing state to the next level. We therefore train one set of measurement and combine layers per Pyramid level, assemble them into PyrInternalModelA/B, and finally test the complete players in a tournament.
The representation rule is the same throughout QSeaBattle:
the game and player wrappers use bits; the trainable models use logits.
A logical bit is represented inside the model by the sign of a logit. In this tutorial we use hard logits -10 and +10.
game bits → player wrapper → gameplay adapter → Pyr model logits → gameplay adapter → game bits
Inside the Pyramid models — including between recursive levels and the shared-resource layers — values remain logits.
from __future__ import annotations
import os
import sys
from pathlib import Path
def change_to_repo_root(marker: str = "src") -> Path:
here = Path.cwd()
for parent in [here] + list(here.parents):
if (parent / marker).is_dir():
os.chdir(parent)
return parent
raise FileNotFoundError(f"Could not find repository root containing {marker!r}")
ROOT = change_to_repo_root("src")
if str(ROOT / "src") not in sys.path:
sys.path.insert(0, str(ROOT / "src"))
print(f"Project root: {ROOT.name}/")
print(f"Working directory: {Path.cwd().name}/")
Project root: QSeaBattle/ Working directory: QSeaBattle/
Imports¶
import numpy as np
import tensorflow as tf
from Q_Sea_Battle.game_env import GameEnv
from Q_Sea_Battle.game_layout import GameLayout
from Q_Sea_Battle.gameplay_adapters import GameplayModelAAdapter, GameplayModelBAdapter
from Q_Sea_Battle.pyr_combine_layer_a import PyrCombineLayerA
from Q_Sea_Battle.pyr_combine_layer_b import PyrCombineLayerB
from Q_Sea_Battle.pyr_dataset_conversion_utilities import (
convert_layer_combine_a,
convert_layer_combine_b,
convert_layer_measure_a,
convert_layer_measure_b,
)
from Q_Sea_Battle.pyr_dataset_generation_utilities import generate_pyr_dataset
from Q_Sea_Battle.pyr_internal_model_a import PyrInternalModelA
from Q_Sea_Battle.pyr_internal_model_b import PyrInternalModelB
from Q_Sea_Battle.pyr_measurement_layer_a import PyrMeasurementLayerA
from Q_Sea_Battle.pyr_measurement_layer_b import PyrMeasurementLayerB
from Q_Sea_Battle.tournament import Tournament
from Q_Sea_Battle.trainable_assisted_players import TrainableAssistedPlayers
tf.get_logger().setLevel("ERROR")
print("TensorFlow:", tf.__version__)
TensorFlow: 2.21.0
1. Choose the Pyramid and Training Settings¶
For a 4×4 field we have n² = 16 positions and four Pyramid levels: 16 → 8 → 4 → 2.
The measurement layers learn the local decisions at each level. The combine layers learn how the current level is reduced to the next one; these parity-like mappings generally need more training.
SMOKE_TEST=True is useful after refactoring because it checks every API and tensor path quickly. For the published result use the full settings.
FIELD_SIZE = 4
COMMS_SIZE = 1
P_RULE = 1.0
BETA = 10.0
ALPHA = 0.3
SEED = 123
SMOKE_TEST = False
if SMOKE_TEST:
DATASET_SIZE = 2_000
BATCH_SIZE = 64
MAX_EPOCHS_MEAS = 3
MAX_EPOCHS_COMB = 5
TARGET_ACCURACY = 0.95
else:
DATASET_SIZE = 50_000
BATCH_SIZE = 256
MAX_EPOCHS_MEAS = 30
MAX_EPOCHS_COMB = 150
TARGET_ACCURACY = 0.9995
tf.random.set_seed(SEED)
np.random.seed(SEED)
layout = GameLayout(
field_size=FIELD_SIZE,
comms_size=COMMS_SIZE,
number_of_games_in_tournament=1000,
channel_noise=0.0,
enemy_probability=0.5,
)
n2 = FIELD_SIZE * FIELD_SIZE
depth = int(np.log2(n2))
levels = [n2 // (2 ** d) for d in range(depth)]
print("Pyramid levels:", levels)
print("Hard-logit magnitude β:", BETA)
Pyramid levels: [16, 8, 4, 2] Hard-logit magnitude β: 10.0
2. Generate the Classical Imitation Data¶
The canonical Pyramid generator produces the known-correct strategy in bits. Before a tensor is presented to a trainable layer, we convert it to the model representation.
With β = 10:
- bit
0→ logit-10 - bit
1→ logit+10
A raw 0/1 tensor should therefore never be passed directly into an internal layer that expects logits.
canonical_ds = generate_pyr_dataset(
n2=n2,
num_games=DATASET_SIZE,
seed=SEED,
validate=True,
)
layer_data_meas_a = convert_layer_measure_a(
canonical_ds,
rep_x="hard_logit",
rep_y="hard_logit",
beta=BETA,
)
layer_data_comb_a = convert_layer_combine_a(
canonical_ds,
rep_field="hard_logit",
rep_outcome="hard_logit",
rep_target="hard_logit",
beta=BETA,
)
layer_data_meas_b = convert_layer_measure_b(
canonical_ds,
rep_x="hard_logit",
rep_y="hard_logit",
beta=BETA,
)
layer_data_comb_b = convert_layer_combine_b(
canonical_ds,
rep_gun="hard_logit",
rep_outcome_b="hard_logit",
rep_comm_in="hard_logit",
rep_gun_next="hard_logit",
rep_comm_next="hard_logit",
beta=BETA,
)
print("Converted imitation data for", depth, "levels.")
Converted imitation data for 4 levels.
3. Train the Learned Layers¶
Targets remain logits. Binary cross-entropy needs binary labels, so the loss converts only the sign of the target logit to 0/1. Accuracy is also measured by sign agreement.
For Bob's combine layer there are two outputs — the reduced gun state and the reduced communication state — so we report both accuracies explicitly rather than relying on Keras-generated history names.
def logits_bce_from_targets(y_true_logits, y_pred_logits):
labels = tf.cast(tf.cast(y_true_logits, tf.float32) >= 0.0, tf.float32)
return tf.reduce_mean(
tf.nn.sigmoid_cross_entropy_with_logits(
labels=labels,
logits=tf.cast(y_pred_logits, tf.float32),
)
)
def sign_accuracy(y_true_logits, y_pred_logits) -> float:
y_true = tf.cast(y_true_logits, tf.float32) >= 0.0
y_pred = tf.cast(y_pred_logits, tf.float32) >= 0.0
return float(tf.reduce_mean(tf.cast(tf.equal(y_true, y_pred), tf.float32)).numpy())
class StopAtAccuracy(tf.keras.callbacks.Callback):
def __init__(self, target=0.9995):
super().__init__()
self.target = float(target)
def on_epoch_end(self, epoch, logs=None):
# Works for single-output models. Multi-output models simply train to max_epochs.
logs = logs or {}
accs = [float(v) for k, v in logs.items() if "sign_acc" in k and isinstance(v, (int, float))]
if accs and min(accs) >= self.target:
self.model.stop_training = True
def train_single_input_layer(layer, x_train, y_train, *, epochs, learning_rate=1e-3):
x_in = tf.keras.Input(shape=x_train.shape[1:], dtype=tf.float32)
y_out = layer(x_in)
model = tf.keras.Model(x_in, y_out)
model.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate),
loss=logits_bce_from_targets,
)
model.fit(
x_train,
y_train,
epochs=epochs,
batch_size=BATCH_SIZE,
verbose=0,
)
pred = layer(tf.convert_to_tensor(x_train, tf.float32), training=False)
return layer, sign_accuracy(y_train, pred)
def train_two_input_layer(layer, x_left, x_right, y_train, *, epochs, learning_rate=1e-3):
left = tf.keras.Input(shape=x_left.shape[1:], dtype=tf.float32)
right = tf.keras.Input(shape=x_right.shape[1:], dtype=tf.float32)
y_out = layer(left, right)
model = tf.keras.Model([left, right], y_out)
model.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate),
loss=logits_bce_from_targets,
)
model.fit(
[x_left, x_right],
y_train,
epochs=epochs,
batch_size=BATCH_SIZE,
verbose=0,
)
pred = layer(
tf.convert_to_tensor(x_left, tf.float32),
tf.convert_to_tensor(x_right, tf.float32),
training=False,
)
return layer, sign_accuracy(y_train, pred)
def train_b_combine_layer(layer, gun, sr, comm, y_gun, y_comm, *, epochs, learning_rate=1e-3):
gun_in = tf.keras.Input(shape=gun.shape[1:], dtype=tf.float32)
sr_in = tf.keras.Input(shape=sr.shape[1:], dtype=tf.float32)
comm_in = tf.keras.Input(shape=comm.shape[1:], dtype=tf.float32)
gun_out, comm_out = layer(gun_in, sr_in, comm_in)
model = tf.keras.Model([gun_in, sr_in, comm_in], [gun_out, comm_out])
model.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate),
loss=[logits_bce_from_targets, logits_bce_from_targets],
)
model.fit(
[gun, sr, comm],
[y_gun, y_comm],
epochs=epochs,
batch_size=BATCH_SIZE,
verbose=0,
)
pred_gun, pred_comm = layer(
tf.convert_to_tensor(gun, tf.float32),
tf.convert_to_tensor(sr, tf.float32),
tf.convert_to_tensor(comm, tf.float32),
training=False,
)
return (
layer,
sign_accuracy(y_gun, pred_gun),
sign_accuracy(y_comm, pred_comm),
)
meas_layers_a = []
comb_layers_a = []
meas_layers_b = []
comb_layers_b = []
for d, L in enumerate(levels):
print(f"\nLevel {d}: L={L}")
x_a, y_a = layer_data_meas_a[d]
layer_ma, acc_ma = train_single_input_layer(
PyrMeasurementLayerA(hidden_units=16),
x_a, y_a,
epochs=MAX_EPOCHS_MEAS,
)
meas_layers_a.append(layer_ma)
print(f" Measure A : {acc_ma:.4f}")
(field_d, sr_a_d), target_a_d = layer_data_comb_a[d]
layer_ca, acc_ca = train_two_input_layer(
PyrCombineLayerA(hidden_units=16),
field_d, sr_a_d, target_a_d,
epochs=MAX_EPOCHS_COMB,
)
comb_layers_a.append(layer_ca)
print(f" Combine A : {acc_ca:.4f}")
x_b, y_b = layer_data_meas_b[d]
layer_mb, acc_mb = train_single_input_layer(
PyrMeasurementLayerB(hidden_units=16),
x_b, y_b,
epochs=MAX_EPOCHS_MEAS,
)
meas_layers_b.append(layer_mb)
print(f" Measure B : {acc_mb:.4f}")
(gun_d, sr_b_d, comm_d), (gun_next_d, comm_next_d) = layer_data_comb_b[d]
layer_cb, acc_g, acc_c = train_b_combine_layer(
PyrCombineLayerB(hidden_units=16),
gun_d, sr_b_d, comm_d,
gun_next_d, comm_next_d,
epochs=MAX_EPOCHS_COMB,
)
comb_layers_b.append(layer_cb)
print(f" Combine B gun : {acc_g:.4f}")
print(f" Combine B comm : {acc_c:.4f}")
print("\nLayer imitation complete.")
Level 0: L=16 Measure A : 1.0000 Combine A : 1.0000 Measure B : 1.0000 Combine B gun : 1.0000 Combine B comm : 0.9887 Level 1: L=8 Measure A : 1.0000 Combine A : 1.0000 Measure B : 1.0000 Combine B gun : 1.0000 Combine B comm : 1.0000 Level 2: L=4 Measure A : 1.0000 Combine A : 1.0000 Measure B : 1.0000 Combine B gun : 1.0000 Combine B comm : 1.0000 Level 3: L=2 Measure A : 1.0000 Combine A : 1.0000 Measure B : 1.0000 Combine B gun : 1.0000 Combine B comm : 1.0000 Layer imitation complete.
4. Assemble the Recursive Pyramid Models¶
Each trained layer now occupies its corresponding Pyramid level. We install these learned components directly into fresh PyrInternalModelA/B instances.
The shared-resource mechanism itself was not trained: it remains the exact PR-assisted correlation layer. During gameplay we use stochastic mode; p_rule=1 gives the ideal correlation.
No intermediate tensor is converted back to bits inside the model.
model_a = PyrInternalModelA(
layout,
sr_mode="stochastic",
p_rule=P_RULE,
beta=BETA,
alpha=ALPHA,
seed=SEED,
measure_layers=meas_layers_a,
combine_layers=comb_layers_a,
)
model_b = PyrInternalModelB(
layout,
sr_mode="stochastic",
p_rule=P_RULE,
beta=BETA,
alpha=ALPHA,
measure_layers=meas_layers_b,
combine_layers=comb_layers_b,
)
print("Model A levels:", len(model_a.measure_layers), "measurement /", len(model_a.combine_layers), "combine")
print("Model B levels:", len(model_b.measure_layers), "measurement /", len(model_b.combine_layers), "combine")
Model A levels: 4 measurement / 4 combine Model B levels: 4 measurement / 4 combine
5. Check One Internal Forward Pass¶
The learned logits do not need to have magnitude exactly 10, but their sign carries the logical value. The shared-resource outcomes and hardened recursive interfaces should remain clearly separated from zero.
This diagnostic is useful after refactoring because an accidental 0/1 tensor at a logit interface can make the complete strategy fail even when individual calls still run.
def describe(name, x):
x = tf.cast(x, tf.float32)
print(
f"{name:24s}",
"shape=", tuple(x.shape),
"min=", f"{float(tf.reduce_min(x)):+.2f}",
"max=", f"{float(tf.reduce_max(x)):+.2f}",
"mean|x|=", f"{float(tf.reduce_mean(tf.abs(x))):.2f}",
)
sample = generate_pyr_dataset(n2=n2, num_games=4, seed=SEED + 1, validate=True)
field_bits = tf.convert_to_tensor(sample["field_bits"][:, 0, :], tf.float32)
gun_bits = tf.convert_to_tensor(sample["gun_bits"][:, 0, :], tf.float32)
field_logits = BETA * (2.0 * field_bits - 1.0)
gun_logits = BETA * (2.0 * gun_bits - 1.0)
comm_a_logits, meas_a_logits, out_a_logits = model_a.compute_with_internal(
field_logits=field_logits,
harden_between_levels=True,
beta_for_hardening=BETA,
training=False,
)
shoot_logit, meas_b_logits, out_b_logits, comms_b_logits, gun_logits_list = model_b.compute_with_internal(
gun_logits=gun_logits,
comm_in_logits=comm_a_logits,
prev_meas_list=meas_a_logits,
prev_out_list=out_a_logits,
harden_between_levels=True,
beta_for_hardening=BETA,
training=False,
)
describe("field -> A", field_logits)
for d in range(depth):
describe(f"A meas level {d}", meas_a_logits[d])
describe(f"A out level {d}", out_a_logits[d])
describe("A communication", comm_a_logits)
describe("gun -> B", gun_logits)
for d in range(depth):
describe(f"B meas level {d}", meas_b_logits[d])
describe(f"B out level {d}", out_b_logits[d])
describe("shoot", shoot_logit)
field -> A shape= (4, 16) min= -10.00 max= +10.00 mean|x|= 10.00 A meas level 0 shape= (4, 8) min= -22.50 max= +22.15 mean|x|= 11.78 A out level 0 shape= (4, 8) min= -10.00 max= +10.00 mean|x|= 10.00 A meas level 1 shape= (4, 4) min= -14.66 max= +12.44 mean|x|= 9.92 A out level 1 shape= (4, 4) min= -10.00 max= +10.00 mean|x|= 10.00 A meas level 2 shape= (4, 2) min= -9.79 max= +9.73 mean|x|= 9.14 A out level 2 shape= (4, 2) min= -10.00 max= +10.00 mean|x|= 10.00 A meas level 3 shape= (4, 1) min= -9.38 max= +11.40 mean|x|= 9.97 A out level 3 shape= (4, 1) min= -10.00 max= +10.00 mean|x|= 10.00 A communication shape= (4, 1) min= +17.33 max= +17.37 mean|x|= 17.35 gun -> B shape= (4, 16) min= -10.00 max= +10.00 mean|x|= 10.00 B meas level 0 shape= (4, 8) min= -24.23 max= +6.74 mean|x|= 12.88 B out level 0 shape= (4, 8) min= -10.00 max= +10.00 mean|x|= 10.00 B meas level 1 shape= (4, 4) min= -13.94 max= +8.75 mean|x|= 10.63 B out level 1 shape= (4, 4) min= -10.00 max= +10.00 mean|x|= 10.00 B meas level 2 shape= (4, 2) min= -11.37 max= +9.59 mean|x|= 10.79 B out level 2 shape= (4, 2) min= -10.00 max= +10.00 mean|x|= 10.00 B meas level 3 shape= (4, 1) min= -12.32 max= +10.88 mean|x|= 11.60 B out level 3 shape= (4, 1) min= -10.00 max= +10.00 mean|x|= 10.00 shoot shape= (4, 1) min= -22.23 max= +17.36 mean|x|= 20.25
6. Run the Complete Players in a Tournament¶
The gameplay adapters are the representation boundary. They receive bit-valued game state from the player wrappers, convert it to logits for the internal Pyramid models, and convert the final outputs back to bits.
Tournament → TrainableAssistedPlayers (bits) → Gameplay adapters → PyrInternalModelA/B (logits)
For p_rule=1 and zero channel noise, the ideal Pyramid strategy wins every game. The trained imitation model should therefore approach a score of 1.0.
adapter_a = GameplayModelAAdapter(
internal_model_a=model_a,
beta=BETA,
harden_between_levels=True,
)
adapter_b = GameplayModelBAdapter(
internal_model_b=model_b,
beta=BETA,
harden_between_levels=True,
)
eval_layout = GameLayout(
field_size=FIELD_SIZE,
comms_size=COMMS_SIZE,
number_of_games_in_tournament=1000,
channel_noise=0.0,
enemy_probability=0.5,
)
env = GameEnv(eval_layout)
players = TrainableAssistedPlayers(
eval_layout,
model_a=adapter_a,
model_b=adapter_b,
)
tournament = Tournament(
game_env=env,
players=players,
game_layout=eval_layout,
)
log = tournament.tournament()
mean_reward, std_err = log.outcome()
print(f"Tournament score: {mean_reward:.4f} ± {std_err:.4f}")
print("Ideal score at p_rule=1.0: 1.0000")
Tournament score: 0.9880 ± 0.0034 Ideal score at p_rule=1.0: 1.0000
7. Save the Bootstrap Weights¶
These imitation-trained models provide a useful known-good initialization for later end-to-end optimization. We save weights rather than serialized models so the checkpoints remain tied to the current QSeaBattle model classes.
models_dir = ROOT / "models"
models_dir.mkdir(parents=True, exist_ok=True)
model_a_path = models_dir / f"pyr_model_a_bootstrap_f{FIELD_SIZE}_r{P_RULE:.2f}.weights.h5"
model_b_path = models_dir / f"pyr_model_b_bootstrap_f{FIELD_SIZE}_r{P_RULE:.2f}.weights.h5"
model_a.save_weights_to(str(model_a_path))
model_b.save_weights_to(str(model_b_path))
print("Saved:", model_a_path.name)
print("Saved:", model_b_path.name)
Saved: pyr_model_a_bootstrap_f4_r1.00.weights.h5 Saved: pyr_model_b_bootstrap_f4_r1.00.weights.h5