Reza2kn's picture
Release complete mixed Q3/Q4 Audio8 with portable CPU, MLX and exact GGUF GPU consumers
c49eca9 verified
Raw History Blame Contribute Delete
4.41 kB
"""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,'<u2').astype(np.uint32)<<16).view(np.float32)
require(np.isfinite(values).all(),'nonfinite BF16')
array=mx.array(values.reshape(r['shape'])).astype(mx.bfloat16);mx.eval(array);item=Dense(array)
else:item=packed_record(stream,r)
tensors[r['name']]=item;progress(dict(name=r['name'],resident_tensor_bytes=item.nbytes))
require(fmt.sha256(path)==m['model_sha256'] and fmt.sha256(manifest_path)==manifest_sha,'native model changed while loading')
for alias,target in m['aliases'].items():tensors[alias]=tensors[target]
return cls(tensors,dict(manifest_sha256=manifest_sha,weights_sha256=m['model_sha256'],source_revision='b4413de154ed6bdef0a4011028b1ebd12aca8152',profile='native_q4_attention_q3_mlp',source_weight_bytes=m['payload_bytes'],native_container=True,expected_config_sha256=fmt.CONFIG_SHA256,layout='mlx_affine_power_of_two_offset',affine_bias_storage='transient_derived_from_scale',refit=False,source_payload_sha_verified=True,production_ready=False))