MachineSex/src/neural/speciation_real.py
Giorgio Gilestro ea051a5f92 E13b/c: harden real-weight speciation — full symmetry group + emergent-divergence null
E13c (the symmetry defense): alignment now runs modulo the FULL
function-preserving unit symmetry group of a ReLU MLP (per-unit positive
rescaling via canonicalise_scale, composed with Re-Basin permutations;
sanity gate recovers a permuted-and-rescaled copy exactly). Verdict: the
full group removes the independent-init barrier (residual 0.001) and
essentially none of the conflict barrier (0.502 -> 0.497) — the residual
is functional, not a missed symmetry (answers arXiv:2606.23607). The
cliff gains a hybrid-fitness readout: merged accuracy 0.97 -> 0.03 with
conflict. Floor proposition drafted (paper/si-notes.md S1): endpoint
invariance + max(eps_A, eps_B) >= mu(S)/2 for any merged model under any
alignment group.

E13b (emergent divergence): pre-registered second reading — with NO
conflicting training signal (disjoint class specialists; rolled-input
conventions), residual is 0.000 at every divergence to t_div=3200, and
the merge RESCUES the forgetting specialists (parents 0.535/0.474 ->
merged 0.955; a sustained Fisher-Muller rescue at zero barrier).
Speciation in real weights requires functional conflict; it does not
emerge from compatible specialisation on shared ancestry. LLM-scale
over-specialisation (cf. 2607.11997) deferred to Phase-3 llm_speciation.

3-panel figure, READMEs, +2 tests (149 green), make mnist wired.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01BkRLcc18rwT2Lysu6PbG7v
2026-09-06 12:35:14 +01:00

205 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""E13 — real-weight model speciation: the merge-compatibility cliff in trained MLPs, with Git Re-Basin.
The real-weight image of E12. Two small MLPs are forked from a shared base and trained independently;
we merge them (weight averaging) and measure the linear-mode-connectivity **barrier** — before and after
**permutation alignment** (:mod:`neural.rebasin`). The barrier alignment *removes* is a coordinate
artefact (Git Re-Basin); the barrier it *cannot* remove — the **residual** — is the true reproductive-
isolation / DobzhanskyMuller signal. A ``condition`` knob sets how the children diverge, which is the
empirical question this experiment answers:
- ``shared`` — both children keep training on the *same* task from the shared fork: they never leave the
basin, so there is ~no barrier at all (trivially mergeable — the low-divergence anchor).
- ``independent`` — same task, but each child trained from its *own random init* (the canonical Git
Re-Basin setting): a large naive barrier that alignment *removes* (residual ≈ 0) — the coordinate
artefact. "Same species, different basis."
- ``conflict`` — the children learn *conflicting* label maps (B's labels cyclically shifted):
genuinely incompatible functions on shared capacity. A large barrier that alignment *cannot* remove
(residual stays high) — true reproductive isolation. "Different species."
- ``disjoint`` (E13b, *emergent* divergence) — child A keeps training only on classes 04, child B only
on 59: no contradiction anywhere (a true DobzhanskyMuller setting — each lineage's changes are
harmless alone). Does a residual barrier *emerge* with divergence, without imposed conflict? And does
the *merged* model first rescue the two forgetting specialists (FisherMuller) then fail (speciation)
as divergence grows — E12's compatible → depression → inviability curve, emergent in real weights?
- ``augment`` (E13b, conventions) — same task and labels, but A trains on images rolled +3 px and B on
images rolled 3 px: representational conventions drift with zero output conflict.
The discriminating metric is the **residual** (barrier after alignment): ~0 for ``shared`` and
``independent`` (compatible — the incompatibility, if any, is coordinate), large for ``conflict``.
Alignment is reported at two levels (E13c): permutation-only (Git Re-Basin, ``residual``) and
**scale-canonicalised + permutation** (``residual_scale``) — the *full* function-preserving unit
symmetry group of a plain ReLU MLP — so the residual cannot be attributed to a symmetry the aligner
missed (cf. arXiv:2606.23607). Merged-model (midpoint) accuracies are recorded alongside the barriers.
Divergence is swept via post-fork training steps ``t_div``. Small no-BatchNorm MLPs on MNIST — the clean
Re-Basin regime. Statistically reproducible (seeded); NumPy/scipy alignment is deterministic.
"""
from __future__ import annotations
from typing import Any, Mapping
import numpy as np
import pandas as pd
from .rebasin import apply_perms, barrier, canonicalise_scale, interpolate, weight_matching
from .train import seed_everything
def _mlp(sizes, device):
import torch.nn as nn
layers = []
for i in range(len(sizes) - 1):
layers.append(nn.Linear(sizes[i], sizes[i + 1]))
if i < len(sizes) - 2:
layers.append(nn.ReLU())
return nn.Sequential(*layers).to(device)
def _get_params(model):
import torch.nn as nn
return [(m.weight.detach().cpu().numpy().copy(), m.bias.detach().cpu().numpy().copy())
for m in model if isinstance(m, nn.Linear)]
def _set_params(model, params, device):
import torch
import torch.nn as nn
it = iter(params)
for m in model:
if isinstance(m, nn.Linear):
W, b = next(it)
m.weight.data = torch.tensor(W, dtype=torch.float32, device=device)
m.bias.data = torch.tensor(b, dtype=torch.float32, device=device)
def _train(model, X, y, steps, lr, batch, rng, device):
import torch
opt = torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9)
lossf = torch.nn.CrossEntropyLoss()
model.train()
for _ in range(steps):
idx = rng.integers(0, len(X), size=batch)
xb = X[idx]; yb = y[idx]
opt.zero_grad(); loss = lossf(model(xb), yb); loss.backward(); opt.step()
def run_speciation_real(cfg: Mapping[str, Any], seed: int) -> pd.DataFrame:
"""Fork-and-merge sweep over (condition x divergence x replicate); returns barrier decompositions."""
import torch
import torchvision
spec = dict(cfg["speciation_real"])
sizes = list(spec.get("sizes", [784, 512, 512, 10]))
conditions = list(spec.get("conditions", ["shared", "independent", "conflict"]))
t_divs = list(spec.get("t_div", [50, 100, 200, 400, 800]))
base_steps = int(spec.get("base_steps", 300))
lr, batch = float(spec.get("lr", 0.05)), int(spec.get("batch", 128))
n_eval = int(spec.get("n_eval", 2000))
reps = int(cfg.get("n_replicates", spec.get("reps", 3)))
device = "cuda" if __import__("torch").cuda.is_available() else "cpu"
root = spec.get("data_root", "data")
tr = torchvision.datasets.MNIST(root, train=True, download=True)
te = torchvision.datasets.MNIST(root, train=False, download=True)
Xtr = (tr.data.float().reshape(-1, 784) / 255.0).to(device); ytr = tr.targets.to(device)
Xte = (te.data.float().reshape(-1, 784) / 255.0)[:n_eval].to(device); yte = te.targets[:n_eval].to(device)
lossf = torch.nn.CrossEntropyLoss()
def make_loss_fn(model):
def loss_fn(params):
_set_params(model, params, device); model.eval()
with torch.no_grad():
out = model(Xte)
L = float(lossf(out, yte)); E = float((out.argmax(1) != yte).float().mean())
return L, E
return loss_fn
def _conflict_labels(frac):
"""B's labels: cyclically shift the first round(frac*10) classes (systematic conflict on them)."""
k = int(round(frac * 10))
if k == 0:
return ytr
sel = np.arange(k); shifted = np.roll(sel, 1)
y2 = ytr.clone()
for c, c2 in zip(sel, shifted):
y2[ytr == int(c)] = int(c2)
return y2
def _rolled(px):
"""Images shifted horizontally by ``px`` pixels (a representational convention; labels intact)."""
return torch.roll(Xtr.reshape(-1, 28, 28), shifts=px, dims=2).reshape(-1, 784)
def subset(cond, child, conflict_frac):
"""(X, y) the child trains on for a given condition."""
if cond == "conflict" and child == 1:
return Xtr, _conflict_labels(conflict_frac) # B learns a conflicting label map
if cond == "disjoint": # emergent DMI: disjoint, compatible tasks
mask = (ytr < 5) if child == 0 else (ytr >= 5)
return Xtr[mask], ytr[mask]
if cond == "augment": # emergent conventions: same task, shifted views
return _rolled(3 if child == 0 else -3), ytr
return Xtr, ytr # shared / independent / conflict-A: normal task
# Two modes: (1) conditions x t_div decomposition; (2) a conflict-fraction isolation cliff.
conflict_fracs = spec.get("conflict_fracs")
if conflict_fracs is not None:
sweep = [("conflict", int(spec.get("t_div_fixed", 800)), float(f)) for f in conflict_fracs]
else:
sweep = [(c, int(t), 1.0) for c in conditions for t in t_divs]
rows: list[dict] = []
for rep in range(reps):
for cond, t_div, conflict_frac in sweep:
ss = seed + 1000 * rep + hash((cond, t_div, conflict_frac)) % 997
seed_everything(np.random.SeedSequence(ss))
base = _mlp(sizes, device)
_train(base, Xtr, ytr, base_steps, lr, batch, np.random.default_rng(ss), device)
base_params = _get_params(base)
children = []
for child in (0, 1):
Xc, yc = subset(cond, child, conflict_frac)
if cond == "independent": # each child from its OWN init (Git Re-Basin regime)
seed_everything(np.random.SeedSequence(ss + 100 * (child + 1)))
m = _mlp(sizes, device)
_train(m, Xc, yc, base_steps + t_div, lr, batch,
np.random.default_rng(ss + 17 * (child + 1)), device)
else: # shared fork, then independent divergence
m = _mlp(sizes, device); _set_params(m, base_params, device)
_train(m, Xc, yc, t_div, lr, batch, np.random.default_rng(ss + 17 * (child + 1)), device)
children.append(_get_params(m))
pA, pB = children
probe = _mlp(sizes, device); loss_fn = make_loss_fn(probe)
b_naive = barrier(pA, pB, loss_fn)
# E13: permutation-only alignment (Git Re-Basin) — the coordinate artefact.
perms = weight_matching(pA, pB, np.random.default_rng(ss + 5))
pB_perm = apply_perms(pB, perms)
b_aligned = barrier(pA, pB_perm, loss_fn)
# E13c: scale-canonicalise both, then match — the FULL ReLU unit symmetry group, so the
# residual cannot be blamed on a symmetry the aligner missed (arXiv:2606.23607).
cA, cB = canonicalise_scale(pA), canonicalise_scale(pB)
perms_c = weight_matching(cA, cB, np.random.default_rng(ss + 5))
cB_al = apply_perms(cB, perms_c)
b_scale = barrier(cA, cB_al, loss_fn)
# E13b money curve: the merged (midpoint) model vs its parents on the full task.
accA, accB = 1 - loss_fn(pA)[1], 1 - loss_fn(pB)[1]
acc_mid_naive = 1 - loss_fn(interpolate(pA, pB, 0.5))[1]
acc_mid_aligned = 1 - loss_fn(interpolate(pA, pB_perm, 0.5))[1]
acc_mid_scale = 1 - loss_fn(interpolate(cA, cB_al, 0.5))[1]
rows.append({
"condition": cond, "t_div": t_div, "conflict_frac": conflict_frac, "replicate": rep,
"barrier_naive": b_naive["error_barrier"],
"barrier_aligned": b_aligned["error_barrier"],
"removable": b_naive["error_barrier"] - b_aligned["error_barrier"],
"residual": b_aligned["error_barrier"],
"barrier_aligned_scale": b_scale["error_barrier"],
"removable_scale": b_naive["error_barrier"] - b_scale["error_barrier"],
"residual_scale": b_scale["error_barrier"],
"loss_barrier_naive": b_naive["loss_barrier"],
"loss_barrier_aligned": b_aligned["loss_barrier"],
"loss_barrier_scale": b_scale["loss_barrier"],
"acc_parent_a": accA, "acc_parent_b": accB,
"acc_merge_naive": acc_mid_naive,
"acc_merge_aligned": acc_mid_aligned,
"acc_merge_scale": acc_mid_scale,
"parent_acc": (accA + accB) / 2.0})
return pd.DataFrame(rows)