File size: 3,729 Bytes
b2ebc95
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
#!/usr/bin/env python3
"""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:
        # Expected here: the base ships a 15-tensor MTP (multi-token-prediction) head that
        # prime-rl's SFT trainer does not train and does not save. The served artifact is A's
        # structure, so B's extras are ignored rather than grafted on — the output must be
        # loadable by exactly the command in SUBMISSION.md.
        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)

    # Load B once, keyed by tensor name — B's shard layout is irrelevant from here on.
    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()