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.

In [9]:
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¶

In [10]:
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.

In [11]:
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.

In [12]:
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.

In [13]:
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),
    )
In [14]:
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.

In [15]:
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.

In [16]:
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.

In [17]:
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.

In [18]:
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