E13: real-weight model speciation — the Git Re-Basin residual confirms E12

The real-weight image of E12, and the answer to the mode-connectivity reviewer.
Small no-BN MLPs on MNIST, forked from a shared base and trained independently,
are weight-averaged; we measure the linear-mode-connectivity barrier before and
after in-house deterministic Git Re-Basin permutation alignment (neural/rebasin.py,
scipy linear_sum_assignment), decomposing it into removable (coordinate artefact)
and residual (reproductive isolation). kind: speciation_real.

Result (3 reps):
- shared (same task, shared fork): no barrier — trivially mergeable.
- independent (same task, different init): naive 0.056, alignment removes 98%
  (residual 0.001) — the incompatibility is a coordinate artefact.
- conflict (conflicting label maps): naive 0.496, alignment removes 0% (residual
  0.496) — genuine reproductive isolation. Because alignment demonstrably works on
  the independent case, the conflict residual is real, not a failure to align.
- Isolation cliff (speciation_real_cliff): residual rises 0.00->0.13->0.19->0.28->
  0.40->0.49 with the fraction of conflicting classes — the real-weight mirror of
  E12's cliff; residual==naive throughout (functional, not coordinate).

rebasin.py sanity-gated (recovers a known permutation exactly). plot_speciation_real.py
(2-panel), +4 pure-NumPy tests (142 green), README with honest positioning vs
Git Re-Basin / Entezari / Frankle / Pari 2024 / Zhou 2026. Wired into make mnist
(needs torchvision).

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Giorgio Gilestro 2026-07-09 00:19:04 +01:00
parent 56f642e7f9
commit 01d87e504f
15 changed files with 620 additions and 3 deletions

View file

@ -188,6 +188,11 @@ def run_and_save(config_path: str | Path) -> Path:
elif kind == "mnist_lineage":
df, extra_manifest = run_mnist_experiment(cfg) # oracle provenance + confusion matrix
grid = [{"label": label, "lineage_cfg": c} for label, c in _expand_mnist(cfg)]
elif kind == "speciation_real":
from .speciation_real import run_speciation_real # E13: real-weight speciation + Git Re-Basin
df = run_speciation_real(cfg, int(cfg["seed"]))
grid = None
extra_manifest = {"layer": "1.5", "tier": "speciation_real"}
else:
raise ValueError(f"unknown neural experiment kind {kind!r}")
save_artifacts(cfg, df, out_dir, extra_libs=_EXTRA_LIBS,

110
src/neural/rebasin.py Normal file
View file

@ -0,0 +1,110 @@
"""Git Re-Basin weight-matching + linear-mode-connectivity barrier for small MLPs (in-house, deterministic).
The real-weight speciation experiment (E13) needs to separate the part of two models' merge-incompatibility
that is a mere *coordinate artefact* (removable by permuting hidden units Ainsworth et al., Git Re-Basin,
arXiv:2209.04836) from the *residual* incompatibility that permutation cannot fix (the true reproductive-
isolation / DobzhanskyMuller signal). This module implements, for a plain MLP with no BatchNorm/residuals
(the clean Re-Basin regime):
- ``weight_matching`` align model B's hidden units to model A by per-layer coordinate descent, each layer
an exact linear-assignment problem (``scipy.optimize.linear_sum_assignment``). Deterministic given the two
weight sets and a seeded layer-visiting order; no data, no gradient so the "coordinate artefact" claim is
not confounded with data.
- ``barrier`` the linear-mode-connectivity loss/error barrier along the interpolation between two models.
Weights are handled as a list of ``(W, b)`` numpy arrays, one per ``nn.Linear`` (``W`` is ``[out, in]``).
Permutations act on the *hidden* layers (outputs of all but the last linear), permuting rows of layer ``k``
(and its bias) and columns of layer ``k+1``.
"""
from __future__ import annotations
import numpy as np
from scipy.optimize import linear_sum_assignment
def weight_matching(params_a: list, params_b: list, rng: np.random.Generator,
max_iter: int = 30) -> list:
"""Permutations aligning B's hidden units to A (Git Re-Basin weight matching).
Args:
params_a, params_b (list[tuple[np.ndarray, np.ndarray]]): ``(W, b)`` per linear layer; ``W`` is
``[out, in]``. Both models must share architecture.
rng (np.random.Generator): seeds the (order-only) layer-visiting schedule deterministic result.
max_iter (int): coordinate-descent sweeps; stops early at a fixed point.
Returns:
list[np.ndarray]: permutation index arrays, one per hidden layer (length = #linear 1).
"""
n_layers = len(params_a)
hidden = n_layers - 1 # permutable layers (not the output)
sizes = [params_a[k][0].shape[0] for k in range(hidden)]
perms = [np.arange(s) for s in sizes] # identity to start
for _ in range(max_iter):
changed = False
for k in rng.permutation(hidden):
Wa_in, _ = params_a[k]
Wb_in, _ = params_b[k]
# incoming weights: B's columns already permuted by the previous hidden layer's perm
if k > 0:
Wb_in = Wb_in[:, perms[k - 1]]
cost = Wa_in @ Wb_in.T # (out_k, out_k): align A row i with B row j
# outgoing weights: layer k+1 columns correspond to these units; B's rows permuted by next perm
Wa_out = params_a[k + 1][0]
Wb_out = params_b[k + 1][0]
if k + 1 < hidden:
Wb_out = Wb_out[perms[k + 1], :]
cost = cost + Wa_out.T @ Wb_out # (in_{k+1}=out_k, out_k)
ri, ci = linear_sum_assignment(cost, maximize=True)
new_perm = ci[np.argsort(ri)]
if not np.array_equal(new_perm, perms[k]):
changed = True
perms[k] = new_perm
if not changed:
break
return perms
def apply_perms(params: list, perms: list) -> list:
"""Return a copy of ``params`` with hidden-unit permutations applied (rows of k, columns of k+1)."""
out = [(W.copy(), b.copy()) for W, b in params]
hidden = len(params) - 1
for k in range(hidden):
p = perms[k]
W, b = out[k]
out[k] = (W[p, :], b[p]) # permute this layer's output units
Wn, bn = out[k + 1]
out[k + 1] = (Wn[:, p], bn) # and the next layer's matching inputs
return out
def interpolate(params_a: list, params_b: list, alpha: float) -> list:
"""Linear interpolation ``(1-alpha)*A + alpha*B`` of two parameter lists."""
return [((1 - alpha) * Wa + alpha * Wb, (1 - alpha) * ba + alpha * bb)
for (Wa, ba), (Wb, bb) in zip(params_a, params_b)]
def barrier(params_a: list, params_b: list, loss_fn, alphas=None) -> dict:
"""Linear-mode-connectivity barrier: worst-case excess loss on the interpolation path above the chord.
Args:
params_a, params_b (list): endpoint parameter lists.
loss_fn (Callable): ``loss_fn(params) -> (loss, error)`` on a fixed eval set.
alphas (Sequence[float] | None): interpolation grid (default 21 points on [0, 1]).
Returns:
dict: ``{loss_barrier, error_barrier, midpoint_loss_barrier}`` (excess over the endpoint chord).
"""
alphas = np.linspace(0, 1, 21) if alphas is None else np.asarray(alphas)
losses, errors = [], []
for a in alphas:
L, E = loss_fn(interpolate(params_a, params_b, float(a)))
losses.append(L); errors.append(E)
losses, errors = np.array(losses), np.array(errors)
chord_L = (1 - alphas) * losses[0] + alphas * losses[-1]
chord_E = (1 - alphas) * errors[0] + alphas * errors[-1]
mid = len(alphas) // 2
return {"loss_barrier": float(np.max(losses - chord_L)),
"error_barrier": float(np.max(errors - chord_E)),
"midpoint_loss_barrier": float(losses[mid] - chord_L[mid])}

View file

@ -0,0 +1,165 @@
"""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."
The discriminating metric is the **residual** (barrier after alignment): ~0 for ``shared`` and
``independent`` (compatible the incompatibility, if any, is coordinate), large for ``conflict``.
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, 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 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
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)
perms = weight_matching(pA, pB, np.random.default_rng(ss + 5))
b_aligned = barrier(pA, apply_perms(pB, perms), loss_fn)
accA, accB = 1 - loss_fn(pA)[1], 1 - loss_fn(pB)[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"],
"loss_barrier_naive": b_naive["loss_barrier"],
"loss_barrier_aligned": b_aligned["loss_barrier"],
"parent_acc": (accA + accB) / 2.0})
return pd.DataFrame(rows)