Ling-3.0-tiny-RKNN / tools /append_mtp_package.py
Sariel00's picture
Keep MTP opt-in: default-off build, isolated experimental package and measurements
3706f1c verified
Raw History Blame Contribute Delete
5.49 kB
#!/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()