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()
|