File size: 4,641 Bytes
da1a4ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
#!/usr/bin/env python3
"""Pre-upload verification of the INT8 package against the upstream download. Read-only on both trees
except for writing SHA256SUMS into the package. Exits non-zero on the first failed check.

  usage: verify_package.py UPSTREAM_DIR PACKAGE_DIR
"""
import hashlib
import json
import os
import sys
from pathlib import Path

import torch
from safetensors import safe_open


def fail(msg):
    sys.exit(f"VERIFY_FAIL {msg}")


def shard_map(d):
    index = json.loads((d / "model.safetensors.index.json").read_text())
    return index["weight_map"]


def main():
    up, pkg = Path(sys.argv[1]), Path(sys.argv[2])

    # 1. connector: bf16 file == upstream fp32 cast to bf16, tensor by tensor
    up_map = shard_map(up / "connector")
    with safe_open(str(pkg / "connector/model.safetensors"), "pt") as new:
        if set(new.keys()) != set(up_map):
            fail("connector tensor set differs from upstream")
        n = 0
        for shard in sorted(set(up_map.values())):
            with safe_open(str(up / "connector" / shard), "pt") as old:
                for key in old.keys():
                    ref = old.get_tensor(key)
                    ref = ref.to(torch.bfloat16) if ref.is_floating_point() else ref
                    got = new.get_tensor(key)
                    if got.dtype != ref.dtype or not torch.equal(got, ref):
                        fail(f"connector {key} != fp32->bf16")
                    n += 1
    print(f"OK connector: {n} tensors equal upstream fp32 -> bf16", flush=True)

    # 2. unchanged components: hardlink (same inode) or identical bytes
    same = 0
    for comp in ("transformer", "vae", "mlp", "scheduler"):
        for f in sorted((up / comp).rglob("*")):
            if f.is_dir():
                continue
            g = pkg / f.relative_to(up)
            if not g.is_file():
                fail(f"missing {g}")
            if os.stat(f).st_ino != os.stat(g).st_ino and f.read_bytes() != g.read_bytes():
                fail(f"{g} differs from upstream")
            same += 1
    if (up / "LICENSE").read_bytes() != (pkg / "LICENSE").read_bytes():
        fail("LICENSE differs from upstream")
    print(f"OK unchanged components: {same} files identical to upstream (+ LICENSE)", flush=True)

    # 3. mllm: copied tensors byte-identical, quantized ones present as int8 + fp32 scale
    manifest = json.loads((pkg / "mllm/int8_manifest.json").read_text())
    quant = set(manifest["quantized_modules"])
    old_map, new_map = shard_map(up / "mllm"), shard_map(pkg / "mllm")
    expect_new = {k for k in old_map if k[: -len(".weight")] not in quant or not k.endswith(".weight")}
    expect_new |= {m + ".weight" for m in quant} | {m + ".scale" for m in quant}
    if set(new_map) != expect_new:
        fail(f"mllm index: {len(set(new_map) ^ expect_new)} names differ from the expected set")
    handles = {}

    def tensor(tree, mapping, key):
        path = str(tree / mapping[key])
        if path not in handles:
            handles[path] = safe_open(path, "pt")
        return handles[path].get_tensor(key)

    copied = quantized = 0
    for key in sorted(old_map):
        module = key[: -len(".weight")] if key.endswith(".weight") else None
        ref = tensor(up / "mllm", old_map, key)
        if module in quant:
            w, s = tensor(pkg / "mllm", new_map, key), tensor(pkg / "mllm", new_map, module + ".scale")
            if w.dtype != torch.int8 or s.dtype != torch.float32 or w.shape != ref.shape or s.shape != (ref.shape[0],):
                fail(f"{key}: int8/scale dtype or shape wrong")
            quantized += 1
        else:
            got = tensor(pkg / "mllm", new_map, key)
            if got.dtype != ref.dtype or not torch.equal(got, ref):
                fail(f"{key}: copied tensor differs from upstream")
            copied += 1
        if len(handles) > 4:
            handles.clear()
    print(f"OK mllm: {copied} tensors byte-identical to upstream, {quantized} quantized (int8 + fp32 scale)", flush=True)

    # 4. sha256 of every file in the package
    lines = []
    for f in sorted(p for p in pkg.rglob("*") if p.is_file() and p.name != "SHA256SUMS" and ".cache" not in p.parts):
        h = hashlib.sha256()
        with open(f, "rb") as fh:
            for chunk in iter(lambda: fh.read(1 << 24), b""):
                h.update(chunk)
        lines.append(f"{h.hexdigest()}  {f.relative_to(pkg).as_posix()}")
    (pkg / "SHA256SUMS").write_text("\n".join(lines) + "\n")
    print(f"OK sha256: {len(lines)} files -> SHA256SUMS", flush=True)
    print("VERIFY_OK", flush=True)


if __name__ == "__main__":
    main()