"""Load the verified native container using the existing packed MLX operators.""" import hashlib,json,struct import numpy as np import mlx.core as mx from . import native_model_format as fmt from .weights import Dense,packed_record,require def load_native(cls,directory,*,progress): manifest_path=directory/'manifest.json';manifest_sha=fmt.sha256(manifest_path) m=json.loads(manifest_path.read_text());path=directory/'model.a8m' require(m['format']=='A8MOD001' and m['status']=='complete','incomplete native export') require(path.stat().st_size==m['model_bytes'] and fmt.sha256(path)==m['model_sha256'],'native file identity') require(m['source_manifest_sha256']=='a068c0532c24b208d95de2b02af4533701f21bdb2436936178a35c8e6d219f45' and m['source_weights_sha256']=='478a045ad7f447d36205d411c93cf5e5847e0879d950c9ee7dcfc4cf8d289694','pinned source lineage') require(fmt.sha256(directory/'config.json')==fmt.CONFIG_SHA256,'native config identity') verified=fmt.verify_stream(path) require(verified==m['readback'],'native schema/footer identity') require(m['aliases']=={'language_model.lm_head.weight':'language_model.model.embed_tokens.weight'},'native tied head') expected=fmt.schema();require(set(m['tensor_sha256'])==set(expected),'native tensor hash inventory') records=[] # Hash every payload before allocating GPU weights. The separate format # verifier has already checked all shapes, names, grids and total length. with path.open('rb') as stream: stream.seek(fmt.HEADER.size) for name,(shape,kind) in expected.items(): n,k,rank,group,bits,zero,maximum,size=fmt.RECORD.unpack(stream.read(fmt.RECORD.size)) dims=list(struct.unpack('<'+'I'*rank,stream.read(rank*4))) require(stream.read(n).decode('ascii')==name and dims==shape and k==kind,'native record changed') offset=stream.tell();h=hashlib.sha256();remaining=size while remaining: block=stream.read(min(fmt.CHUNK,remaining));require(block,'truncated payload');h.update(block);remaining-=len(block) require(h.hexdigest()==m['tensor_sha256'][name],'native tensor hash mismatch') r=dict(name=name,shape=shape,offset=offset,bytes=size,precision='original',encoding='original',source_dtype='BF16') if kind: rows,cols=shape;codes=rows*cols*bits//8 r.update(precision='q3_full8' if bits==3 else 'q4',encoding='uniform_lsb_full8' if bits==3 else 'uniform_lsb',group_size=group,padded_cols=cols,row_bytes=cols*bits//8,codes_bytes=codes,scales_offset=offset+codes,scales_bytes=rows*(cols//64)*2,scales_dtype='F16',scales_shape=[rows,cols//64]) if bits==3:r['grid']=dict(bits=3,zero_point=4,max_code=7,padding_code=4,signed_min=-4,signed_max=3,inference_permutation_required=False) records.append(r) wanted={f'language_model.model.layers.{i}.mlp.{n}.weight' for i in range(36) for n in ('gate_proj','up_proj','down_proj')} require({r['name'] for r in records if r['precision']=='q3_full8'}==wanted,'full MLP Q3 coverage') tensors={} for r in records: if r['precision']=='original': stream.seek(r['offset']);data=stream.read(r['bytes']);require(len(data)==r['bytes'],'truncated BF16') values=(np.frombuffer(data,'