"""
ZERO WORLD v2 — GPU edition (PyTorch). Canonical build.

Every human decision is in bias_ledger.md. If it is not there, it is not here.
Device: Apple Metal (mps) > CUDA > CPU, chosen automatically.

  python3 zeroworld_gpu.py --goal off            # Experiment 1: no purpose
  python3 zeroworld_gpu.py --goal on             # Experiment 2: predict + guess the hidden rule
  add  --push https://.../metrics --push_token X # publish metrics to the public dashboard

v2 adds (all in the ledger):
  - agents sense how many other bodies are in each of the 3x3 cells
  - Exp 2: a heritable bit gives read access to the energy field as it was
    20 steps ago (recorded history). Half the founders have it. We track who
    predicts better and who finds the rule.
  - prediction is scored at horizons 1, 5 and 20 steps (only horizon 1 is rewarded)
  - richer novelty ruler: action shares + lifespan + crowding + mobility + eating
"""
import argparse, json, os, time, urllib.request
import numpy as np
import torch


def pick_device():
    if torch.backends.mps.is_available():
        return torch.device("mps")
    if torch.cuda.is_available():
        return torch.device("cuda")
    return torch.device("cpu")


class ZeroWorld:
    N_LOCAL = 9
    N_IN = 9 + 1 + 9 + 9          # local energy, own energy, other bodies, history
    N_ACT, N_PRED, N_RULE = 5, 9, 8
    N_OUT = N_ACT + 3 * N_PRED + N_RULE   # actions, pred h1, pred h5, pred h20, rule bits
    H_MAX = 16
    HIST_BACK = 20
    HORIZONS = (1, 5, 20)

    def __init__(self, size, pop_cap, seed, goal, qrng_file, outdir, device, leak=0.0, flip_every=0, trickle=0.0, plastic=False):
        self.d = device
        # Experiment 3 physics (ledger 9, 10, and the line-2 trickle). All zero = Experiments 1 and 2 unchanged.
        self.LEAK, self.FLIP_EVERY, self.TRICKLE, self.PLASTIC = leak, flip_every, trickle, plastic
        self.rule_flips, self.rule_flip_step = 0, 0
        self.rule_best_bits, self.rule_held_since = 0, None   # true running maximum, and the step since which the majority guess has been 5 of 5
        self.size, self.pop_cap, self.goal, self.outdir = size, pop_cap, goal, outdir
        os.makedirs(outdir, exist_ok=True)
        if qrng_file and os.path.exists(qrng_file):
            raw = np.fromfile(qrng_file, dtype=np.uint8)
            if raw.size >= 8:
                seed = int(np.frombuffer(raw[:8].tobytes(), dtype=np.uint64)[0] % (2**62))
        self.g = torch.Generator(device="cpu").manual_seed(seed)
        self.step_i = 0

        # ---- world (ledger 1, 2) ------------------------------------------
        rng = np.random.default_rng(seed)
        self.RULE = int(rng.integers(1, 256))
        self.rule_bits_true = torch.tensor([(self.RULE >> b) & 1 for b in range(5)])
        self.energy = torch.zeros((size, size), device=device)
        c = size // 2
        self.energy[c - 4:c + 5, c - 4:c + 5] = 1.0           # the single seed event
        self.E_IN = pop_cap * 0.01 * 0.5                      # energy entering per step
        self.E_MAX = 3.0
        self.COST_LIVE, self.COST_MOVE, self.COST_NEURON = 0.01, 0.01, 0.0005
        self.BIRTH_AT, self.CHILD_SHARE = 1.5, 0.5
        self.hist = torch.zeros((self.HIST_BACK, size, size), device=device)  # ring buffer of past fields
        self.bodies = torch.zeros((size, size), device=device)                 # bodies per cell

        # ---- agents (ledger 3, 4, 5) ----------------------------------------
        P = pop_cap
        n0 = min(P, max(2000, P // 10))
        self.alive = torch.zeros(P, dtype=torch.bool, device=device); self.alive[:n0] = True
        self.pos = self.r_int(0, size, (P, 2))
        self.pos0 = self.pos.clone()
        self.e = torch.full((P,), 0.8, device=device)
        self.born = torch.zeros(P, dtype=torch.int64, device=device)
        self.ident = torch.arange(P, dtype=torch.int64, device=device)
        self.next_id = P
        self.W1 = self.r_norm((P, self.N_IN, self.H_MAX)) * 0.5
        self.W2 = self.r_norm((P, self.H_MAX, self.N_OUT)) * 0.5
        self.hmask = torch.zeros((P, self.H_MAX), device=device); self.hmask[:, :2] = 1.0
        self.mut = torch.full((P,), 0.05, device=device)
        self.hist_gene = torch.zeros(P, dtype=torch.bool, device=device)
        self.hist_gene[:n0 // 2] = True                        # half the founders can read history
        if self.PLASTIC:                                        # ledger line 11: plasticity genes, half the founders at zero
            self.eta1 = torch.zeros((P, self.N_IN, self.H_MAX), device=device)
            self.eta2 = torch.zeros((P, self.H_MAX, self.N_OUT), device=device)
            half = torch.arange(P, device=device) % 2 == 1
            self.eta1[half] = self.r_norm((int(half.sum()), self.N_IN, self.H_MAX)) * 0.01
            self.eta2[half] = self.r_norm((int(half.sum()), self.H_MAX, self.N_OUT)) * 0.01
            self.G1, self.G2 = self.W1.clone(), self.W2.clone()  # genes: what is inherited. W1/W2 become the living weights.
        # life records for the novelty ruler
        self.acts = torch.zeros((P, self.N_ACT), device=device)
        self.crowd = torch.zeros(P, device=device)
        self.eats = torch.zeros(P, device=device)
        # Exp 2 scoring
        self.pred_err = torch.zeros((P, 3), device=device)    # EMA per horizon
        self.rule_bits = torch.zeros((P, self.N_RULE), device=device)
        self.pq = {h: torch.zeros((P, h, self.N_PRED), device=device) for h in self.HORIZONS}   # prediction queues
        self.pqpos = {h: torch.zeros((P, h, 2), dtype=torch.int64, device=device) for h in self.HORIZONS}
        self.pqok = {h: torch.zeros((P, h), dtype=torch.bool, device=device) for h in self.HORIZONS}

        # ---- logs ------------------------------------------------------------
        self.lineage_buf = []
        self.lineage_path = os.path.join(outdir, "lineage.bin")   # int64 rows: step, parent, child, hidden, mut*1e6, hist_gene
        self.archive = torch.zeros((0, 9), device=device)
        self.archive_log = []
        self.NOVELTY_T = 0.25
        self.deaths_total = 0

    # random helpers: CPU generator -> device (portable across mps/cuda/cpu)
    def r_int(self, lo, hi, shape):
        return torch.randint(lo, hi, shape, generator=self.g).to(self.d)
    def r_norm(self, shape):
        return torch.randn(shape, generator=self.g).to(self.d)
    def r_uni(self, shape):
        return torch.rand(shape, generator=self.g).to(self.d)

    def patch(self, field, pos):
        y, x = pos[:, 0], pos[:, 1]
        cols = []
        for dy in (-1, 0, 1):
            for dx in (-1, 0, 1):
                cols.append(field[(y + dy) % self.size, (x + dx) % self.size])
        return torch.stack(cols, dim=1)

    def regrow(self):
        on = (self.energy > 0.05).to(torch.int64)
        nb = torch.roll(on, 1, 0) + torch.roll(on, -1, 0) + torch.roll(on, 1, 1) + torch.roll(on, -1, 1)
        bits = torch.tensor([(self.RULE >> k) & 1 for k in range(5)], device=self.d, dtype=torch.bool)
        grow = bits[nb]
        k = int(grow.sum().item())
        e_rule = self.E_IN * (1.0 - self.TRICKLE)
        if self.TRICKLE > 0:   # ledger line 2, from Experiment 3: a trickle to every cell so the world cannot starve itself black
            self.energy = torch.clamp(self.energy + self.E_IN * self.TRICKLE / self.energy.numel(), max=1.0)
        if k:
            self.energy = torch.where(grow, torch.clamp(self.energy + e_rule / k, max=1.0), self.energy)

    def step(self):
        if self.FLIP_EVERY and self.step_i > 0 and self.step_i % self.FLIP_EVERY == 0:   # ledger line 10: the rule moves
            b = int(torch.randint(0, 5, (1,), generator=self.g))
            self.RULE ^= (1 << b)
            self.rule_bits_true = torch.tensor([(self.RULE >> k) & 1 for k in range(5)])
            self.rule_flips += 1; self.rule_flip_step = self.step_i
            with open(os.path.join(self.outdir, "hidden_rule.txt"), "a") as f:
                f.write(f"step={self.step_i} flipped bit {b} -> RULE={self.RULE} bits={[(self.RULE>>k)&1 for k in range(5)]}\n")
        idx = torch.nonzero(self.alive, as_tuple=False).squeeze(1)
        n = idx.numel()
        if n == 0:
            return
        pos = self.pos[idx]

        # ---- sense ------------------------------------------------------------
        local = self.patch(self.energy, pos)
        others = self.patch(self.bodies, pos)
        others[:, 4] -= 1.0                                   # do not count yourself
        past = self.patch(self.hist[self.step_i % self.HIST_BACK], pos) * self.hist_gene[idx, None].float()
        inp = torch.cat([local, self.e[idx, None], others, past], dim=1)

        # ---- think -----------------------------------------------------------
        h = torch.tanh(torch.bmm(inp[:, None, :], self.W1[idx]).squeeze(1)) * self.hmask[idx]
        out = torch.bmm(h[:, None, :], self.W2[idx]).squeeze(1)
        act = out[:, :self.N_ACT].argmax(1)
        self.acts[idx, act] += 1
        if self.PLASTIC:   # ledger line 11: weights move by eta * pre * post; no target, no error
            self.W1[idx] += self.eta1[idx] * (inp[:, :, None] * h[:, None, :])
            self.W2[idx] += self.eta2[idx] * (h[:, :, None] * torch.tanh(out)[:, None, :])
            self.W1[idx].clamp_(-5, 5); self.W2[idx].clamp_(-5, 5)
        self.crowd[idx] += others.sum(1)

        # ---- move ------------------------------------------------------------
        dy = (act == 2).to(torch.int64) - (act == 1).to(torch.int64)
        dx = (act == 4).to(torch.int64) - (act == 3).to(torch.int64)
        ny = (pos[:, 0] + dy) % self.size
        nx = (pos[:, 1] + dx) % self.size
        self.pos[idx, 0] = ny; self.pos[idx, 1] = nx

        # ---- eat: one body per cell wins -------------------------------------
        flat = ny * self.size + nx
        order = torch.argsort(flat)
        sf = flat[order]
        first = torch.ones_like(sf, dtype=torch.bool)
        first[1:] = sf[1:] != sf[:-1]
        win = order[first]
        eflat = self.energy.view(-1)
        take = eflat[flat[win]]
        eflat[flat[win]] = 0.0
        self.e[idx[win]] = torch.clamp(self.e[idx[win]] + take, max=self.E_MAX)
        self.eats[idx[win]] += (take > 0.01).float()

        nh = self.hmask[idx].sum(1)
        self.e[idx] -= self.COST_LIVE + self.COST_MOVE * (act != 0).float() + self.COST_NEURON * nh
        if self.LEAK > 0:      # ledger line 9: stored energy leaks
            self.e[idx] *= (1.0 - self.LEAK)

        # ---- world update ------------------------------------------------------
        self.hist[self.step_i % self.HIST_BACK] = self.energy   # record before regrow = field "20 steps ago" when read
        self.regrow()

        # ---- Experiment 2: prediction + rule guess (ledger 6) ------------------
        if self.goal:
            newpos = self.pos[idx]
            for hi, hz in enumerate(self.HORIZONS):
                slot = self.step_i % hz
                # score the prediction made hz steps ago for the cell we were on then
                ok = self.pqok[hz][idx, slot]
                if bool(ok.any()):
                    sel = idx[ok]
                    actual = self.patch(self.energy, self.pqpos[hz][sel, slot])
                    err = ((self.pq[hz][sel, slot] - actual) ** 2).mean(1)
                    self.pred_err[sel, hi] = 0.9 * self.pred_err[sel, hi] + 0.1 * err
                    if hz == 1:
                        # the only reward in the universe; kept below COST_LIVE on purpose
                        self.e[sel] += 0.008 * torch.clamp(1.0 - err, 0, 1)
                # store the new prediction
                a0 = self.N_ACT + hi * self.N_PRED
                self.pq[hz][idx, slot] = out[:, a0:a0 + self.N_PRED]
                self.pqpos[hz][idx, slot] = newpos
                self.pqok[hz][idx, slot] = True
            self.rule_bits[idx] = 0.95 * self.rule_bits[idx] + 0.05 * (out[:, -self.N_RULE:] > 0).float()

        # ---- death -----------------------------------------------------------
        dead = idx[self.e[idx] <= 0]
        if dead.numel():
            self.archive_behaviours(dead)
            self.alive[dead] = False
            self.deaths_total += dead.numel()

        # ---- birth (ledger 4) --------------------------------------------------
        parents = idx[(self.e[idx] >= self.BIRTH_AT) & self.alive[idx]]
        free = torch.nonzero(~self.alive, as_tuple=False).squeeze(1)
        k = min(parents.numel(), free.numel())
        if k and k < parents.numel():
            parents = parents[torch.randperm(parents.numel(), generator=self.g).to(self.d)]
        if k:
            p, c = parents[:k], free[:k]
            self.e[c] = self.e[p] * self.CHILD_SHARE
            self.e[p] -= self.e[c]
            self.pos[c] = (self.pos[p] + self.r_int(-1, 2, (k, 2))) % self.size
            self.pos0[c] = self.pos[c]
            m = self.mut[p]
            if self.PLASTIC:
                self.G1[c] = self.G1[p] + self.r_norm(self.G1[p].shape) * m[:, None, None]
                self.G2[c] = self.G2[p] + self.r_norm(self.G2[p].shape) * m[:, None, None]
                self.W1[c], self.W2[c] = self.G1[c].clone(), self.G2[c].clone()        # born with the genes, not the parent's lived weights
                self.eta1[c] = self.eta1[p] + self.r_norm(self.eta1[p].shape) * m[:, None, None] * 0.01
                self.eta2[c] = self.eta2[p] + self.r_norm(self.eta2[p].shape) * m[:, None, None] * 0.01
            else:
                self.W1[c] = self.W1[p] + self.r_norm(self.W1[p].shape) * m[:, None, None]
                self.W2[c] = self.W2[p] + self.r_norm(self.W2[p].shape) * m[:, None, None]
            flip = self.r_uni((k,)) < m
            j = self.r_int(0, self.H_MAX, (k,))
            hm = self.hmask[p].clone()
            rows = torch.nonzero(flip, as_tuple=False).squeeze(1)
            hm[rows, j[rows]] = 1 - hm[rows, j[rows]]
            hm[hm.sum(1) == 0, 0] = 1
            self.hmask[c] = hm
            self.mut[c] = torch.clamp(m * torch.exp(self.r_norm((k,)) * 0.2), 1e-4, 1.0)
            self.hist_gene[c] = self.hist_gene[p] ^ (self.r_uni((k,)) < m)   # history access is heritable, can flip
            self.alive[c] = True
            self.born[c] = self.step_i
            self.acts[c] = 0; self.crowd[c] = 0; self.eats[c] = 0
            self.pred_err[c] = 0; self.rule_bits[c] = 0
            for hz in self.HORIZONS:
                self.pqok[hz][c] = False
            new_ids = torch.arange(self.next_id, self.next_id + k, device=self.d)
            self.ident[c] = new_ids
            self.next_id += k
            self.lineage_buf.append(torch.stack([
                torch.full((k,), self.step_i, dtype=torch.int64, device=self.d),
                self.ident[p], new_ids, hm.sum(1).to(torch.int64),
                (self.mut[c] * 1_000_000).to(torch.int64), self.hist_gene[c].to(torch.int64)], dim=1).cpu())

        # ---- body map for next step's senses ----------------------------------
        self.bodies.zero_()
        aidx = torch.nonzero(self.alive, as_tuple=False).squeeze(1)
        self.bodies.view(-1).index_add_(0, self.pos[aidx, 0] * self.size + self.pos[aidx, 1],
                                        torch.ones(aidx.numel(), device=self.d))
        self.step_i += 1

    def archive_behaviours(self, dead):
        a = self.acts[dead]
        tot = a.sum(1, keepdim=True)
        keep = tot[:, 0] >= 20
        if not bool(keep.any()):
            return
        d = dead[keep]
        steps = tot[keep][:, 0]
        a = a[keep] / steps[:, None]
        life = torch.log1p((self.step_i - self.born[d]).float())[:, None] / 10.0
        crowd = torch.clamp(self.crowd[d] / steps, 0, 8)[:, None] / 8.0
        disp = (self.pos[d] - self.pos0[d]).abs().float()
        disp = torch.minimum(disp, self.size - disp).sum(1) / (self.size / 2)   # wrapped net displacement 0..2
        mobility = torch.clamp(disp, 0, 2)[:, None] / 2.0
        eat = (self.eats[d] / steps)[:, None]
        desc = torch.cat([a, life, crowd, mobility, eat], dim=1)              # 9 numbers
        if self.archive.shape[0]:
            dist = torch.cdist(desc, self.archive).min(1).values
            desc = desc[dist > self.NOVELTY_T]
        if desc.shape[0] == 0:
            return
        desc = desc[:64]
        chosen = [desc[0]]
        for row in desc[1:]:
            if torch.cdist(row[None], torch.stack(chosen)).min() > self.NOVELTY_T:
                chosen.append(row)
        self.archive = torch.cat([self.archive, torch.stack(chosen)])
        self.archive_log.append((self.step_i, int(self.archive.shape[0])))

    def snapshot(self, n=256):
        """Downsampled picture of the universe for the public page: energy (mean) and bodies (sum), 0..255."""
        import base64, torch.nn.functional as F
        k = max(1, self.size // n)
        e = F.avg_pool2d(self.energy[None, None], k)[0, 0]
        b = F.avg_pool2d(self.bodies[None, None], k)[0, 0]          # mean bodies per cell, 0..~3
        e8 = (e.clamp(0, 1) * 255).to(torch.uint8).cpu().numpy().tobytes()
        b8 = (b.clamp(0, 1) * 255).to(torch.uint8).cpu().numpy().tobytes()  # 255 = one body in every cell
        return {"n": e.shape[0], "energy": base64.b64encode(e8).decode(), "bodies": base64.b64encode(b8).decode(),
                "seed_at": [self.size // 2 // k, self.size // 2 // k]}

    def flush_lineage(self):
        if self.lineage_buf:
            arr = torch.cat(self.lineage_buf).numpy().astype(np.int64)
            with open(self.lineage_path, "ab") as f:
                arr.tofile(f)
            self.lineage_buf = []

    def metrics(self):
        idx = torch.nonzero(self.alive, as_tuple=False).squeeze(1)
        n = idx.numel()
        m = {"step": self.step_i, "goal": self.goal, "population": int(n),
             "mean_energy": float(self.e[idx].mean()) if n else 0.0,
             "world_energy": float(self.energy.sum()),
             "mean_hidden": float(self.hmask[idx].sum(1).mean()) if n else 0.0,
             "mean_mut": float(self.mut[idx].mean()) if n else 0.0,
             "archive_size": int(self.archive.shape[0]),
             "births_total": int(self.next_id - self.pop_cap),
             "deaths_total": int(self.deaths_total),
             "hist_share": float(self.hist_gene[idx].float().mean()) if n else 0.0,
             "rule_flips": self.rule_flips, "rule_flip_step": self.rule_flip_step}
        if self.PLASTIC and n:
            mag = (self.eta1[idx].abs().mean((1, 2)) + self.eta2[idx].abs().mean((1, 2)))
            pl = mag > 1e-4
            m["plastic_share"] = float(pl.float().mean())
            m["mean_plasticity"] = float(mag.mean())
            if self.goal and bool(pl.any()) and bool((~pl).any()):
                m["pred_err_h1_plastic"] = float(self.pred_err[idx][pl, 0].mean())
                m["pred_err_h1_no_plastic"] = float(self.pred_err[idx][~pl, 0].mean())
        if self.goal and n:
            pe = self.pred_err[idx]
            m["pred_err_h1"], m["pred_err_h5"], m["pred_err_h20"] = [float(v) for v in pe.mean(0)]
            hg = self.hist_gene[idx]
            if bool(hg.any()) and bool((~hg).any()):
                m["pred_err_h1_with_history"] = float(pe[hg, 0].mean())
                m["pred_err_h1_no_history"] = float(pe[~hg, 0].mean())
                gw = (self.rule_bits[idx][hg].mean(0)[:5].cpu() > 0.5).to(torch.int64)
                gn = (self.rule_bits[idx][~hg].mean(0)[:5].cpu() > 0.5).to(torch.int64)
                m["rule_bits_with_history"] = int((gw == self.rule_bits_true).sum())
                m["rule_bits_no_history"] = int((gn == self.rule_bits_true).sum())
            guess = (self.rule_bits[idx].mean(0)[:5].cpu() > 0.5).to(torch.int64)
            m["rule_match_bits"] = int((guess == self.rule_bits_true).sum())
            m["rule_match_of"] = 5
            self.rule_best_bits = max(self.rule_best_bits, m["rule_match_bits"])
            if m["rule_match_bits"] == 5:
                self.rule_held_since = self.step_i if self.rule_held_since is None else self.rule_held_since
            else:
                self.rule_held_since = None
            m["rule_best_bits"] = self.rule_best_bits
            m["rule_held_steps"] = (self.step_i - self.rule_held_since) if self.rule_held_since is not None else 0
        return m


def push(url, token, payload):
    try:
        req = urllib.request.Request(url, data=json.dumps(payload).encode(), method="POST",
                                     headers={"Content-Type": "application/json", "Authorization": f"Bearer {token}"})
        urllib.request.urlopen(req, timeout=10).read()
    except Exception as ex:  # never let the public page stop the world
        print("push failed:", ex, flush=True)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--goal", choices=["on", "off"], default="off")
    ap.add_argument("--steps", type=int, default=1_000_000)
    ap.add_argument("--size", type=int, default=512)
    ap.add_argument("--pop", type=int, default=100_000)
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--qrng", default=None)
    ap.add_argument("--out", default=None)
    ap.add_argument("--log_every", type=int, default=200)
    ap.add_argument("--device", default=None)
    ap.add_argument("--push", default=None, help="public dashboard endpoint")
    ap.add_argument("--push_token", default=os.environ.get("ZEROWORLD_TOKEN"))
    ap.add_argument("--push_every", type=int, default=3, help="push every N logs")
    ap.add_argument("--run_key", default=None, help="feed key; default goal_<on|off>")
    ap.add_argument("--leak", type=float, default=0.0, help="Exp 3: fraction of body energy lost per step (ledger 9)")
    ap.add_argument("--flip_every", type=int, default=0, help="Exp 3: flip one hidden rule bit every N steps (ledger 10)")
    ap.add_argument("--trickle", type=float, default=0.0, help="Exp 3: share of E_IN spread over every cell (ledger 2)")
    ap.add_argument("--plastic", action="store_true", help="Exp 4: heritable Hebbian plasticity per connection (ledger 11)")
    a = ap.parse_args()
    dev = torch.device(a.device) if a.device else pick_device()
    out = a.out or (f"run_{a.run_key}" if a.run_key else f"run_goal_{a.goal}")
    print("device:", dev, "grid:", a.size, "pop cap:", a.pop, flush=True)
    w = ZeroWorld(a.size, a.pop, a.seed, a.goal == "on", a.qrng, out, dev, a.leak, a.flip_every, a.trickle, a.plastic)
    run_key = a.run_key or f"goal_{a.goal}"
    with open(os.path.join(out, "hidden_rule.txt"), "w") as f:
        f.write(f"RULE={w.RULE} bits={[(w.RULE>>b)&1 for b in range(5)]}\n")
    hist, t0 = [], time.time()
    start_iso = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())

    def record(force_push=False, final=False):
        """Log a row; push on cadence, or always when forced (extinction, final step)."""
        m = w.metrics(); m["sec"] = round(time.time() - t0, 1)
        m["steps_per_sec"] = round(w.step_i / max(m["sec"], 1e-6), 1)
        m["grid"], m["pop_cap"], m["device"] = a.size, a.pop, str(dev)
        m["final"] = final
        hist.append(m)
        w.flush_lineage()
        with open(os.path.join(out, "metrics.json"), "w") as f:
            json.dump({"latest": m, "history": hist[::max(1, len(hist) // 2000)]}, f)
        with open(os.path.join(out, "archive.json"), "w") as f:
            json.dump({"growth": w.archive_log, "descriptors": w.archive.cpu().tolist()}, f)
        n_logged = len(hist) - 1
        if a.push and (force_push or n_logged % a.push_every == 0):
            pub = {"latest": m, "history": hist[::max(1, len(hist) // 400)], "snap": w.snapshot(),
                   "hidden_rule_bits": [(w.RULE >> b) & 1 for b in range(5)],
                   "archive": w.archive.cpu().tolist()[-600:], "started_at": start_iso}
            push(a.push, a.push_token, {"run": run_key, **pub})
        print(json.dumps(m), flush=True)
        return m

    for s in range(a.steps):
        w.step()
        if s % a.log_every == 0:
            m = record()
            if m["population"] == 0:
                record(force_push=True, final=True)   # the extinction row always reaches the public feed
                print("EXTINCT at step", w.step_i); break
    else:
        # the loop ran to the end: step_i == a.steps. This row is the one the verdicts judge.
        record(force_push=True, final=True)
    w.flush_lineage()


if __name__ == "__main__":
    main()
