fsi-anomaly / tests /test_merges.py
FerrellSyntheticIntelligence's picture
backup all: 100 files (batch)
76b78ee verified
Raw
History Blame Contribute Delete
1.54 kB
"""Regression tests for post-training merges (tiny-model-posttrain).
Covers the two measured 2026-08-13 merge bugs on the 50M 16k line:
- base pretrain checkpoints carry mtp_heads.* keys that folded
post-training checkpoints lack (must intersect keys, never KeyError)
- trim_delta flattened its mask before indexing the tensor (IndexError)
"""
import torch
from train.ties_merge import ties_merge, trim_delta
def test_trim_delta_preserves_top_fraction_shape():
delta = torch.randn(8, 5)
out = trim_delta(delta, keep=0.2)
assert out.shape == delta.shape
nonzero = (out != 0).sum().item()
assert nonzero > 0
assert nonzero <= delta.numel() # top-20% per tensor, never more
def test_ties_merge_ignores_missing_task_keys():
base = {"w1": torch.randn(4, 4), "mtp_heads.0.weight": torch.randn(4, 4)}
task = {"w1": torch.randn(4, 4)} # folded ckpt: no mtp keys
out = ties_merge(base, [task, task], keep=0.5)
assert "w1" in out
assert "mtp_heads.0.weight" not in out
def test_soup_taskarith_intersect_keys():
from train.parallel_merges import model_soup, task_arithmetic
base = {"w1": torch.randn(4, 4), "mtp_heads.0.weight": torch.randn(4, 4)}
t1 = {"w1": torch.randn(4, 4)}
t2 = {"w1": torch.randn(4, 4)}
soup = model_soup([t1, t2])
assert set(soup.keys()) == {"w1"}
ta = task_arithmetic(base, [t1, t2], lam=0.5)
assert set(ta.keys()) == {"w1"}
assert torch.allclose(ta["w1"], base["w1"] + 0.5 * ((t1["w1"] - base["w1"]) + (t2["w1"] - base["w1"])))