"""Parallel merge recipes for post-training checkpoints (tiny-model-posttrain). Runs model soup (arXiv 2203.05482) and task arithmetic (arXiv 2212.04089) on the SAME base, producing one checkpoint per recipe. TIES is handled by train/ties_merge.py (arXiv 2306.01708). After merging, battery-eval every candidate and keep the best (LFM2 2511.23404 ยง4.4: parallel apply -> eval -> select). Never naive-average adapters; these recipes operate on folded full-weight checkpoints where delta = task_ckpt - base is the true task vector. Usage: .venv/bin/python train/parallel_merges.py \ --base ckpt/hybrid50m_v16k_pretrain/model_5000.pt \ --tasks ckpt/hybrid50m_v25_lora/best.pt ckpt/hybrid50m_v25_dpo/model_final.pt \ --out-dir ckpt/hybrid50m_v25_merges \ --lambda-ta 0.5 """ import argparse from pathlib import Path import torch from model.config import TinyLiquidConfig from model.tiny_liquid import TinyLiquid from model.utils import latest_ckpt def load_sd(path, tag): sd = torch.load(str(path), map_location="cpu", weights_only=False) print(f" {tag}: {path} step={sd.get('step', '?')} tag={sd.get('tag', '-')}", flush=True) return sd def model_soup(task_sds): """Simple average of task weights (all tasks trained from the same base).""" soup = {} keys = [k for k in task_sds[0] if task_sds[0][k].is_floating_point() and all(k in sd for sd in task_sds[1:])] for k in keys: soup[k] = torch.stack([sd[k].float() for sd in task_sds]).mean(dim=0) return soup def task_arithmetic(base_sd, task_sds, lam): """base + lam * sum(task_i - base).""" merged = {} keys = [k for k in base_sd if base_sd[k].is_floating_point() and all(k in sd for sd in task_sds)] for k in keys: base = base_sd[k].float() delta = torch.zeros_like(base) for sd in task_sds: delta += sd[k].float() - base merged[k] = base + lam * delta return merged def main(): ap = argparse.ArgumentParser() ap.add_argument("--base", required=True) ap.add_argument("--tasks", nargs="+", required=True) ap.add_argument("--out-dir", required=True) ap.add_argument("--lambda-ta", type=float, default=0.5) ap.add_argument("--threads", type=int, default=4) args = ap.parse_args() assert len(args.tasks) >= 2, "merges need >= 2 task checkpoints" torch.set_num_threads(args.threads) base_path = Path(args.base) base_ckpt = base_path if base_path.is_file() else latest_ckpt(args.base) base = load_sd(base_ckpt, "base") base_sd = base["model"] task_sds = [] for i, t in enumerate(args.tasks): tp = Path(t) tp = tp if tp.is_file() else latest_ckpt(t) task_sds.append(load_sd(tp, f"task{i}")["model"]) recipes = { "soup": model_soup(task_sds), f"taskarith_l{args.lambda_ta}".replace(".", "p"): task_arithmetic(base_sd, task_sds, args.lambda_ta), } out_dir = Path(args.out_dir) out_dir.mkdir(parents=True, exist_ok=True) for name, merged in recipes.items(): cfg = TinyLiquidConfig(**base["config"]) cfg.mtp_heads = 0 model = TinyLiquid(cfg) missing, unexpected = model.load_state_dict(merged, strict=False) if missing or unexpected: print(f"[{name}] ignored {len(missing)} missing / {len(unexpected)} unexpected keys", flush=True) out = out_dir / f"{name}.pt" torch.save({"config": base["config"], "model": merged, "step": 0, "best_val": base.get("best_val", float("inf")), "tag": f"{name}-{len(task_sds)}tasks"}, out) print("saved", out, flush=True) if __name__ == "__main__": main()