Ling-3.0-tiny-RKNN / tools /pack_selective_model.py
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
6.45 kB
#!/usr/bin/env python3
"""Merge the tested selective candidate into one standalone package, no requantization.
W8 source: BF16 nonrouted matrices except KDA input. Calibrated W4: KDA input.
Routed W4 and all other payloads remain byte-identical to the original package.
Requires only Python's standard library. Existing outputs are never overwritten.
"""
import argparse
import hashlib
import json
import mmap
import zlib
from pathlib import Path
from pack_model import HEADER, ENTRY, align, pad_to
W8 = 'kda_out,mla_in,mla_out,shared,dense,head'
REVISION = 'e3a47d5b986e7141b6efd62597d598ebb392060d'
OFFICIAL_REVISION = '65a6d1d71e01f73ba01e572992bbd69ea92c865f'
class Package:
def __init__(self, path):
self.file = path.open('rb')
self.data = mmap.mmap(self.file.fileno(),0,access=mmap.ACCESS_READ)
self.header = list(HEADER.unpack_from(self.data))
h = self.header
assert h[0] == b'L3RKNN1\0' and h[12] == len(self.data)
checksum = h[14]
h[14] = 0
assert zlib.crc32(HEADER.pack(*h)) & 0xffffffff == checksum
h[14] = checksum
self.entries = {}
for i in range(h[4]):
e = list(ENTRY.unpack_from(self.data,h[7]+i*ENTRY.size))
name = self.data[h[8]+e[0]:h[8]+e[0]+e[1]].decode()
assert name not in self.entries
assert e[12]+e[13] <= len(self.data) and e[15] == 0
self.entries[name] = e
def selected_source(base):
if '.mlp.experts.' in base: return 'original'
if '.attention.' in base and (int(base.split('.')[2])+1)%4 and base.endswith('.qkvfgb'):
return 'calibrated'
return 'official'
def verify(path):
p = Package(path)
for name,e in p.entries.items():
digest = hashlib.sha256()
for offset in range(e[12],e[12]+e[13],8*1024*1024):
digest.update(p.data[offset:min(offset+8*1024*1024,e[12]+e[13])])
if digest.digest() != e[22]: raise ValueError('output checksum mismatch: '+name)
return dict(tensors=len(p.entries),bytes=len(p.data),header_crc32=p.header[14],all_payload_sha256_valid=True)
def main():
p = argparse.ArgumentParser(description=__doc__)
for key in ('original','calibrated','official','output'):
p.add_argument('--'+key,type=Path,required=True)
a = p.parse_args()
partial = a.output.with_suffix(a.output.suffix+'.building')
if a.output.exists() or partial.exists(): raise FileExistsError('output or partial already exists')
packs = {key:Package(getattr(a,key)) for key in ('original','calibrated','official')}
for key,pack in packs.items():
assert pack.header[13].hex() == (OFFICIAL_REVISION if key == 'official' else REVISION)
assert bool(pack.header[6]&0x100) == (key == 'official')
choices = {}
counts = dict(original_w4=0,calibrated_w4=0,w8=0)
for name,e in packs['original'].entries.items():
if e[4] != 3: continue
base = name.removesuffix('.weight')
source = selected_source(base)
other = packs[source].entries[name]
assert other[3] == 2 and other[8:10] == e[8:10]
if source == 'official':
assert other[2] == 1 and other[5] == 0 and other[6] == 0
counts['w8'] += 1
else:
assert other[2:12] == e[2:12]
counts[source+'_w4'] += 1
choices[base] = source
assert counts == dict(original_w4=5888,calibrated_w4=18,w8=91),counts
recipe = dict(format='selective_w4_w8_v1',source_revision=REVISION,
official_source_revision=OFFICIAL_REVISION,w8_families=W8,calibrated_w4_families='kda_in',
linear_counts=counts,shared_stage_default=True,execution='mixed_W4A8_W8A8',
description='Exact merge of tested 2026-09-11 candidate; no further weight quantization',
w8_storage='BF16 source, same per-row W8 construction as tested bridge')
metadata = (json.dumps(recipe,sort_keys=True,indent=2)+'\n').encode()
selected = []
for name,e in packs['original'].entries.items():
base,suffix = name.rsplit('.',1) if '.' in name else (name,'')
source = choices.get(base,'original') if suffix in ('weight','scales','correction') else 'original'
if source == 'official' and suffix in ('scales','correction'): continue
selected.append((name,list(packs[source].entries[name]),source))
e = list(ENTRY.unpack(bytes(ENTRY.size)))
e[3]=1;e[6]=3;e[8]=len(metadata);e[13]=len(metadata);e[22]=hashlib.sha256(metadata).digest()
selected.append(('precision.recipe',e,None))
strings = bytearray()
for name,e,source in selected:
raw = name.encode();e[0]=len(strings);e[1]=len(raw);strings.extend(raw)
header = list(packs['original'].header)
header[4]=len(selected);header[6]=(header[6]&~0x100)|0x200
header[7]=HEADER.size;header[8]=HEADER.size+len(selected)*ENTRY.size;header[9]=len(strings)
header[10]=align(header[8]+len(strings));header[14]=0
a.output.parent.mkdir(parents=True,exist_ok=True)
with partial.open('xb') as out:
out.write(bytes(header[10]))
for i,(name,e,source) in enumerate(selected):
pad_to(out,align(out.tell()))
old_offset,old_bytes=e[12:14]
e[12]=out.tell();e[14]=e[15]=0
digest=hashlib.sha256()
if source is None:
out.write(metadata);digest.update(metadata)
else:
for offset in range(old_offset,old_offset+old_bytes,8*1024*1024):
raw=packs[source].data[offset:min(offset+8*1024*1024,old_offset+old_bytes)]
digest.update(raw);out.write(raw)
if digest.digest()!=e[22]: raise ValueError('source checksum mismatch: '+name)
if (i+1)%3000==0: print(json.dumps(dict(copied=i+1,total=len(selected))),flush=True)
pad_to(out,align(out.tell()))
header[12]=out.tell();header[11]=header[12]-header[10]
header[14]=zlib.crc32(HEADER.pack(*header))&0xffffffff
out.seek(0);out.write(HEADER.pack(*header))
for name,e,source in selected:out.write(ENTRY.pack(*e))
out.write(strings)
out.flush()
result=verify(partial)
partial.rename(a.output)
result.update(recipe=recipe,output=str(a.output),source_payloads_verified=True)
a.output.with_suffix('.merge.json').write_text(json.dumps(result,indent=2)+'\n')
print(json.dumps(result),flush=True)
if __name__=='__main__':main()