"""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)