fsi-anomaly / train /parallel_merges.py
FerrellSyntheticIntelligence's picture
backup all: 37 files (final)
97c39f2 verified
Raw
History Blame Contribute Delete
3.74 kB
"""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()