"""Tests for E13's Git Re-Basin weight-matching + barrier (pure NumPy, always runnable).""" from __future__ import annotations import numpy as np import pytest from neural.rebasin import apply_perms, barrier, interpolate, weight_matching def _mlp(sizes, rng): return [(rng.standard_normal((sizes[i + 1], sizes[i])), rng.standard_normal(sizes[i + 1])) for i in range(len(sizes) - 1)] def _forward(params, X): h = X for i, (W, b) in enumerate(params): h = h @ W.T + b if i < len(params) - 1: h = np.maximum(h, 0) return h def test_weight_matching_recovers_a_known_permutation(): # A random-but-functionally-identical permuted copy is in a different "basis"; weight matching must # recover the permutation, so realigning makes the copy functionally identical to the original again. rng = np.random.default_rng(0) A = _mlp([8, 16, 16, 3], rng) p1, p2 = rng.permutation(16), rng.permutation(16) B = apply_perms(A, [p1, p2]) # functionally identical to A, permuted basis X = rng.standard_normal((32, 8)) assert np.allclose(_forward(A, X), _forward(B, X)) # permutation preserves the function perms = weight_matching(A, B, np.random.default_rng(1)) B_realigned = apply_perms(B, perms) assert np.allclose(_forward(A, X), _forward(B_realigned, X), atol=1e-6) # recovered -> function matches def test_weight_matching_is_deterministic(): rng = np.random.default_rng(2) A, B = _mlp([6, 10, 4], rng), _mlp([6, 10, 4], rng) p1 = weight_matching(A, B, np.random.default_rng(3)) p2 = weight_matching(A, B, np.random.default_rng(3)) assert all(np.array_equal(a, b) for a, b in zip(p1, p2)) # deterministic given inputs + seed def test_interpolate_endpoints_and_barrier_zero_for_identical(): rng = np.random.default_rng(4) A = _mlp([5, 8, 2], rng) B = _mlp([5, 8, 2], rng) assert np.allclose(interpolate(A, B, 0.0)[0][0], A[0][0]) assert np.allclose(interpolate(A, B, 1.0)[0][0], B[0][0]) X = rng.standard_normal((20, 5)); tgt = rng.standard_normal((20, 2)) def loss_fn(P): L = float(((_forward(P, X) - tgt) ** 2).mean()); return L, L b = barrier(A, A, loss_fn) # a model with itself: no barrier assert b["loss_barrier"] == pytest.approx(0.0, abs=1e-9) def test_apply_perms_preserves_function(): rng = np.random.default_rng(5) A = _mlp([4, 7, 7, 3], rng) perms = [rng.permutation(7), rng.permutation(7)] X = rng.standard_normal((16, 4)) assert np.allclose(_forward(A, X), _forward(apply_perms(A, perms), X))