"""Convert verified BF16 components to a native MLX checkpoint, with round-trip checks.""" import argparse import gc import json import shutil import time from pathlib import Path import mlx.core as mx import mlx.nn as nn from mlx.utils import tree_flatten from mlx_vlm.models.qwen_image.weights import load_transformer, load_text_encoder, load_vae from scripts.common import (SOURCE, REVISION, RUNTIME_REVISION, MODIFICATION, NOTICE, eligible, sha256, write_json) LOADERS = dict(transformer=load_transformer, text_encoder=load_text_encoder, vae=load_vae) def save_component(model, folder, config, limit=2_000_000_000): folder.mkdir(parents=True) weights = dict(tree_flatten(model.parameters())) chunks, current, size = [], {}, 0 for name, value in weights.items(): if current and size + value.nbytes > limit: chunks.append(current); current, size = {}, 0 current[name] = value; size += value.nbytes if current: chunks.append(current) index = {'metadata': {'total_size': sum(v.nbytes for v in weights.values()), 'modification_notice': MODIFICATION}, 'weight_map': {}} for i, chunk in enumerate(chunks,1): name = f'model-{i:05d}-of-{len(chunks):05d}.safetensors' mx.save_safetensors(str(folder/name), chunk, metadata={'format':'mlx', 'modification_notice':MODIFICATION, 'source_revision':REVISION}) index['weight_map'].update({key:name for key in chunk}) config['mlx_format'] = True config['_modification_notice'] = MODIFICATION write_json(folder/'config.json', config) write_json(folder/'model.safetensors.index.json', index) return weights def main(): ap=argparse.ArgumentParser() ap.add_argument('--source',type=Path,default=SOURCE) ap.add_argument('--output',type=Path,required=True) ap.add_argument('--bits',type=int,choices=(4,6,8,16),default=8) ap.add_argument('--audit',type=Path,default=Path('artifacts/source-audit.json')) args=ap.parse_args() if args.output.exists(): raise FileExistsError(f'Refusing to overwrite {args.output}') audit=json.loads(args.audit.read_text()) if audit['revision'] != REVISION or Path(audit['source']).resolve()!=args.source.resolve(): raise ValueError('Source audit does not match') # Audit SHA256 was computed in this workspace; reject changes since then. for row in audit['files']: p=args.source/row['path'] if p.stat().st_size != row['size'] or p.stat().st_mtime > args.audit.stat().st_mtime: raise ValueError(f'Source changed since audit: {p}') args.output.mkdir(parents=True) report=dict(name='Image21-MLX',source_model='Qwen/Qwen-Image-2.1',source_revision=REVISION, runtime_revision=RUNTIME_REVISION,method='MLX affine weight-only' if args.bits<16 else 'MLX BF16 layout conversion',bits=args.bits, group_size=64,activation_dtype='bfloat16',vae_dtype='float32',components={}, source_audit_sha256=sha256(args.audit),status='in_progress') write_json(args.output/'conversion.json',report) started=time.perf_counter() for component, loader in LOADERS.items(): print(f'Loading {component}',flush=True) model=loader(args.source) quantized=[] if component!='vae' and args.bits<16: def predicate(path, module): yes=isinstance(module,nn.Linear) and eligible(component,path+'.weight',module.weight.shape) if yes: quantized.append(path) return yes nn.quantize(model,bits=args.bits,group_size=64,mode='affine',class_predicate=predicate) if not quantized: raise RuntimeError(f'No quantized modules: {component}') mx.eval(model.parameters()) config=json.loads((args.source/component/'config.json').read_text()) if quantized: config['quantization']=dict(bits=args.bits,group_size=64,mode='affine') expected=save_component(model,args.output/component,config) print(f'Saved {component}; verifying fresh loader',flush=True) restored=loader(args.output) actual=dict(tree_flatten(restored.parameters())) if set(expected)!=set(actual): raise ValueError(f'Reload keys differ: {component}') for key,value in expected.items(): other=actual[key] if value.dtype!=other.dtype or value.shape!=other.shape or not mx.array_equal(value,other).item(): raise ValueError(f'Reload mismatch: {component}/{key}') report['components'][component]=dict(quantized_modules=quantized, tensor_bytes=sum(v.nbytes for v in expected.values()), tensors=len(expected),exact_roundtrip=True) write_json(args.output/'conversion.json',report) del model,restored,actual,expected gc.collect(); mx.clear_cache() for sub in ('processor','scheduler'): shutil.copytree(args.source/sub,args.output/sub, ignore=shutil.ignore_patterns('._*','.cache','*.lock')) for name in ('model_index.json','LICENSE'): shutil.copy2(args.source/name,args.output/name) (args.output/'Notice').write_text(NOTICE+'\n\nBuilt with Qwen\n'+MODIFICATION+'\n') (args.output/'CHANGES.md').write_text('# Modifications\n\n'+MODIFICATION+'\n\n' 'VAE convolution layout is transposed without reducing FP32 precision. ' 'The entire vision tower, token embeddings, language head, norms, ' 'transformer input/output projections, timestep embedding and modulation remain floating point.\n') report.update(status='converted_and_roundtrip_verified',seconds=time.perf_counter()-started) write_json(args.output/'conversion.json',report) print(json.dumps(report,indent=2),flush=True) if __name__=='__main__': main()