File size: 6,450 Bytes
3fd1a35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
135
136
137
138
#!/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()