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