| |
| """Interpolate two checkpoints in weight space: out = alpha*A + (1-alpha)*B. |
| |
| Usage: interpolate.py OUT CKPT_A CKPT_B [alpha=0.5] |
| |
| Why this exists alongside `soup.py`: `soup.py` averages checkpoints *from the same training run* |
| and assumes they share an index and shard filenames. That assumption does not hold here — the base |
| model ships shards named `model.safetensors-0000N-of-00004.safetensors` while prime-rl's trainer |
| writes `model-0000N-of-00004.safetensors`, so `soup.py` dies with `FileNotFoundError` on the base. |
| This script reads each checkpoint's *own* `model.safetensors.index.json` and matches by tensor |
| name, so it does not care how either side lays its shards out. |
| |
| It also does a real weighted interpolation rather than an equal average, and it verifies that both |
| checkpoints expose exactly the same tensor names and shapes before writing anything — silently |
| averaging a mismatched pair would produce a file that loads and generates garbage. |
| |
| Output carries A's shard layout and A's config/tokenizer/processor files. |
| """ |
|
|
| import json |
| import os |
| import shutil |
| import sys |
|
|
| import torch |
| from safetensors.torch import load_file, save_file |
|
|
| AUX = ( |
| "config.json", "generation_config.json", "tokenizer.json", "tokenizer_config.json", |
| "chat_template.jinja", "merges.txt", "vocab.json", |
| "preprocessor_config.json", "video_preprocessor_config.json", |
| ) |
|
|
|
|
| def index_of(ckpt): |
| """-> {tensor_name: shard_filename}""" |
| with open(os.path.join(ckpt, "model.safetensors.index.json")) as f: |
| return json.load(f)["weight_map"] |
|
|
|
|
| def main(): |
| if len(sys.argv) < 4: |
| sys.exit(__doc__) |
| out, a, b = sys.argv[1], sys.argv[2], sys.argv[3] |
| alpha = float(sys.argv[4]) if len(sys.argv) > 4 else 0.5 |
|
|
| ia, ib = index_of(a), index_of(b) |
| missing = set(ia) - set(ib) |
| if missing: |
| sys.exit(f"{len(missing)} tensors in A are absent from B, e.g. {sorted(missing)[:3]}") |
| extra = set(ib) - set(ia) |
| if extra: |
| |
| |
| |
| |
| print(f"note: {len(extra)} tensors exist only in B and are ignored " |
| f"(e.g. {sorted(extra)[:2]})", flush=True) |
| print(f"{len(ia)} tensors; out = {alpha}*A + {1 - alpha}*B", flush=True) |
|
|
| os.makedirs(out, exist_ok=True) |
|
|
| |
| bt = {} |
| for shard in sorted(set(ib.values())): |
| bt.update(load_file(os.path.join(b, shard))) |
| print(f"loaded B ({len(bt)} tensors)", flush=True) |
|
|
| for shard in sorted(set(ia.values())): |
| sd = load_file(os.path.join(a, shard)) |
| merged = {} |
| for name, ta in sd.items(): |
| tb = bt[name] |
| if ta.shape != tb.shape: |
| sys.exit(f"shape mismatch for {name}: {tuple(ta.shape)} vs {tuple(tb.shape)}") |
| merged[name] = (alpha * ta.to(torch.float32) |
| + (1 - alpha) * tb.to(torch.float32)).to(ta.dtype) |
| save_file(merged, os.path.join(out, shard), metadata={"format": "pt"}) |
| print(f"wrote {shard} ({len(merged)} tensors)", flush=True) |
| del sd, merged |
|
|
| shutil.copy(os.path.join(a, "model.safetensors.index.json"), out) |
| for f in AUX: |
| src = os.path.join(a, f) |
| if os.path.exists(src): |
| shutil.copy(src, out) |
| print("done:", out, "->", sorted(os.listdir(out)), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|