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
205 lines
11 KiB
Python
205 lines
11 KiB
Python
"""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 / Dobzhansky–Muller 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 0–4, child B only
|
||
on 5–9: no contradiction anywhere (a true Dobzhansky–Muller 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 (Fisher–Muller) 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)
|