File size: 5,970 Bytes
4f03424
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
"""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()