Image21-MLX-8bit / scripts /convert.py
ixim's picture
Add files using upload-large-folder tool
4f03424 verified
Raw History Blame Contribute Delete
5.97 kB
"""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()