File size: 5,487 Bytes
3706f1c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
#!/usr/bin/env python3
"""Append quantized layer-24 tensors without requantizing any trunk weight."""
import argparse
import hashlib
import json
import os
from pathlib import Path
import zlib

from pack_model import HEADER, ENTRY, align, pad_to


def index(path):
    with path.open("rb") as stream:
        header = list(HEADER.unpack(stream.read(HEADER.size)))
        if (header[0] != b"L3RKNN1\0" or header[1:4] != [1, HEADER.size, 0x01020304]
                or header[5] != ENTRY.size or header[12] != path.stat().st_size):
            raise ValueError("invalid package header: " + str(path))
        crc = header[14]
        header[14] = 0
        if zlib.crc32(HEADER.pack(*header)) & 0xFFFFFFFF != crc:
            raise ValueError("header CRC mismatch")
        header[14] = crc
        stream.seek(header[8])
        names = stream.read(header[9])
        stream.seek(header[7])
        result = []
        for _ in range(header[4]):
            entry = list(ENTRY.unpack(stream.read(ENTRY.size)))
            name = names[entry[0]:entry[0] + entry[1]].decode("utf-8")
            if not name or entry[3] > 4:
                raise ValueError("invalid tensor entry")
            for offset, size in ((entry[12], entry[13]), (entry[14], entry[15])):
                if size and (offset < header[10] or offset + size > header[12]):
                    raise ValueError("out-of-range tensor: " + name)
            result.append((path, name, entry))
    if len({name for _, name, _ in result}) != len(result):
        raise ValueError("duplicate tensor names")
    return header, result


def copy_range(source, output, offset, size):
    source.seek(offset)
    digest = hashlib.sha256()
    remaining = size
    while remaining:
        chunk = source.read(min(remaining, 8 * 1024 * 1024))
        if not chunk:
            raise IOError("short tensor payload")
        output.write(chunk)
        digest.update(chunk)
        remaining -= len(chunk)
    return digest.digest()


def file_hash(path):
    with path.open("rb") as stream:
        return hashlib.file_digest(stream, "sha256").hexdigest()


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--base", type=Path, required=True)
    parser.add_argument("--mtp-package", type=Path, required=True)
    parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()
    if args.output.exists():
        parser.error("refusing to overwrite an existing package")
    header, trunk = index(args.base)
    other_header, other = index(args.mtp_package)
    if header[15:35] != other_header[15:35]:
        parser.error("MTP architecture differs from trunk")
    mtp = [row for row in other if row[1].startswith("model.layers.24.")]
    if len(mtp) != 798 or any(row[1].startswith("model.layers.24.") for row in trunk):
        parser.error("expected 798 MTP tensors and a trunk without MTP")
    tensors = trunk + mtp
    names = bytearray()
    for _, name, entry in tensors:
        entry[0] = len(names)
        encoded = name.encode("utf-8")
        entry[1] = len(encoded)
        names.extend(encoded)
    header[4] = len(tensors)
    header[7] = HEADER.size
    header[8] = HEADER.size + len(tensors) * ENTRY.size
    header[9] = len(names)
    header[10] = align(header[8] + len(names))
    args.output.parent.mkdir(parents=True, exist_ok=True)
    sources = {}
    entries = []
    try:
        with args.output.open("xb") as output:
            output.write(bytes(HEADER.size + len(tensors) * ENTRY.size))
            output.write(names)
            pad_to(output, header[10])
            for path, name, entry in tensors:
                if path not in sources:
                    sources[path] = path.open("rb")
                source = sources[path]
                data_offset = align(output.tell())
                pad_to(output, data_offset)
                if copy_range(source, output, entry[12], entry[13]) != entry[22]:
                    raise ValueError("tensor SHA256 mismatch: " + name)
                aux_offset = align(output.tell()) if entry[15] else 0
                if entry[15]:
                    pad_to(output, aux_offset)
                    copy_range(source, output, entry[14], entry[15])
                entry[12], entry[14] = data_offset, aux_offset
                entries.append(ENTRY.pack(*entry))
            header[12] = align(output.tell())
            pad_to(output, header[12])
            header[11] = header[12] - header[10]
            header[14] = 0
            header[14] = zlib.crc32(HEADER.pack(*header)) & 0xFFFFFFFF
            output.seek(0)
            output.write(HEADER.pack(*header))
            output.writelines(entries)
            output.flush()
            os.fsync(output.fileno())
    finally:
        for source in sources.values():
            source.close()
    record = {"base": args.base.name, "base_sha256": file_hash(args.base),
              "mtp_package": args.mtp_package.name, "mtp_source_sha256": file_hash(args.mtp_package),
              "output": args.output.name, "output_sha256": file_hash(args.output),
              "bytes": header[12], "flags": header[6], "trunk_tensors_unchanged": len(trunk),
              "appended_mtp_tensors": len(mtp), "recomputed_quantization": False}
    args.output.with_suffix(".manifest.json").write_text(json.dumps(record, indent=2) + "\n")
    print(json.dumps(record, indent=2), flush=True)


if __name__ == "__main__":
    main()