Instructions to use ixim/Image21-MLX-8bit with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use ixim/Image21-MLX-8bit with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir Image21-MLX-8bit ixim/Image21-MLX-8bit
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download scripts/convert.py from ixim/Image21-MLX-8bit: direct link, hf CLI and curl.
- Browser
- Download file 5.97 kB
-
https://huggingface.co/ixim/Image21-MLX-8bit/resolve/main/scripts/convert.py
- Command line
-
hf download hf://ixim/Image21-MLX-8bit/scripts/convert.py
-
curl -L -o convert.py https://huggingface.co/ixim/Image21-MLX-8bit/resolve/main/scripts/convert.py
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() | |