File size: 2,987 Bytes
ce94aa4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Memory-efficient verification. Streams tensors one at a time via
safetensors mmap instead of loading whole models into RAM.
Verifies: merged LM == A, merged vision == B.
"""

import hashlib
import argparse
import json
from pathlib import Path
import numpy as np
from safetensors import safe_open


def build_key_index(model_dir: Path) -> dict:
    """Map each tensor name -> the shard file that contains it."""
    idx = {}
    for f in sorted(model_dir.glob("*.safetensors")):
        with safe_open(f, framework="numpy") as sf:  # metadata only
            for k in sf.keys():
                idx[k] = f
    return idx


def get_tensor_np(index: dict, key: str):
    """Fetch a single tensor as float32-normalized numpy for hashing."""
    f = index[key]
    with safe_open(f, framework="pt") as sf:   # pt handles bf16
        t = sf.get_tensor(key)
    # torch tensor -> float32 numpy (lossless for bf16; passthrough others)
    import torch
    if t.dtype == torch.bfloat16:
        t = t.to(torch.float32)
    return t.numpy()


def h(np_arr) -> str:
    return hashlib.sha256(np_arr.tobytes()).hexdigest()[:16]


def check(name, merged_idx, ref_idx, prefix):
    keys = sorted(k for k in merged_idx if k.startswith(prefix))
    mismatches = []
    for i, k in enumerate(keys):
        if k not in ref_idx:
            mismatches.append((k, "missing in reference"))
            continue
        hm = h(get_tensor_np(merged_idx, k))
        hr = h(get_tensor_np(ref_idx, k))
        if hm != hr:
            mismatches.append((k, "hash differs"))
        if (i + 1) % 100 == 0:
            print(f"  {name}: checked {i+1}/{len(keys)}…")
    print(f"\n=== {name} ===")
    print(f"Checked {len(keys)}  mismatches: {len(mismatches)}")
    for k, why in mismatches[:30]:
        print("  !!", k, "-", why)
    return not mismatches


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("-a", "--finetune", required=True)
    ap.add_argument("-b", "--base", required=True)
    ap.add_argument("-m", "--merged", required=True)
    ap.add_argument("--only", choices=["lm", "vision", "both"],
                    default="both")
    args = ap.parse_args()

    print("Indexing (metadata only, no tensor loads)…")
    a_idx = build_key_index(Path(args.finetune))
    b_idx = build_key_index(Path(args.base))
    m_idx = build_key_index(Path(args.merged))

    lm_ok = vis_ok = True
    if args.only in ("lm", "both"):
        lm_ok = check("LM (merged vs A)", m_idx, a_idx, "language_model.")
    if args.only in ("vision", "both"):
        vis_ok = check("VISION (merged vs B)", m_idx, b_idx, "vision_tower.")

    print("\n=== SUMMARY ===")
    if args.only in ("lm", "both"):
        print(f"LM  == A: {'PASS' if lm_ok else 'FAIL'}")
    if args.only in ("vision", "both"):
        print(f"VIS == B: {'PASS' if vis_ok else 'FAIL'}")
    print("✅ VERIFIED" if (lm_ok and vis_ok) else "❌ MISMATCH")


if __name__ == "__main__":
    main()