File size: 7,755 Bytes
9116984
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""Quantize Qwen-Image-2.1 linear weights to portable SDNQ UINT4."""
import argparse
import gc
import inspect
import json
import shutil
from collections import Counter
from pathlib import Path

import torch

from scripts.integrity import verify_files

REVISION = 'b3179ad355be050328e483a9dfdd9e60cd62adfa'
DIFFUSERS_COMMIT = '80c7ed262aeffbeb43ef13ae04baeb9b84515a69'
TRANSFORMER_SKIP = ['proj_out', 'img_in', 'txt_in', 'time_text_embed', 'modulation', 'norm_out']
TEXT_SKIP = ['lm_head', 'visual']
NOTICE = (
    'Qwen is licensed under the Qwen RESEARCH LICENSE AGREEMENT, Copyright (c) 2026 '
    'Hangzhou Tongyi Laboratory Technology Co., Ltd. All Rights Reserved.\n\n'
    'Built with Qwen\n'
    'Independent derivative: Image21-INT4, by ixim / iximbox.\n'
    'Modified files: transformer and text_encoder weight shards, shard indexes and '
    'config.json files. Eligible linear weights converted to SDNQ UINT4 with SVD rank 32; '
    'floating-point exceptions are documented in component quantization reports.\n'
    'VAE, scheduler and processor retained from the pinned upstream snapshot.\n')


def sdnq_config(skip, device):
    import os
    os.environ.setdefault('DIFFUSERS_SDNQ_TRANSFORMERS', '1')
    from sdnq import SDNQConfig
    quant_device = torch.device(device)
    requested = dict(weights_dtype='uint4', use_svd=True, svd_rank=32,
                     modules_to_not_convert=list(skip), use_quantized_matmul=False,
                     quantization_device=quant_device, return_device=torch.device('cpu'))
    parameters = inspect.signature(SDNQConfig).parameters
    accepts_keywords = any(item.kind == inspect.Parameter.VAR_KEYWORD for item in parameters.values())
    kwargs = requested if accepts_keywords else {key: value for key, value in requested.items() if key in parameters}
    missing = [key for key in ('weights_dtype', 'use_svd') if key not in kwargs]
    if missing:
        raise RuntimeError(f'SDNQConfig cannot express {missing}')
    return SDNQConfig(**kwargs)


def _progress(component):
    import sdnq.quantizer as quantizer
    original = quantizer.sdnq_quantize_layer
    seen = {'count': 0}

    def wrapped(*args, **kwargs):
        seen['count'] += 1
        name = kwargs.get('param_name', args[3] if len(args) > 3 else '')
        if seen['count'] == 1 or seen['count'] % 10 == 0:
            print(f'  {component}: quantized {seen["count"]} layers, latest {name}', flush=True)
        return original(*args, **kwargs)

    quantizer.sdnq_quantize_layer = wrapped
    return original


def quantize_component(source, target, component, device):
    source, target = Path(source), Path(target)
    out = target / component
    if out.exists() and any(out.iterdir()):
        raise FileExistsError(f'Refusing to overwrite existing component: {out}')
    if component == 'transformer':
        from diffusers import QwenImage21Transformer2DModel
        cls, skip = QwenImage21Transformer2DModel, TRANSFORMER_SKIP
    elif component == 'text_encoder':
        from transformers import Qwen3VLForConditionalGeneration
        cls, skip = Qwen3VLForConditionalGeneration, TEXT_SKIP
    else:
        raise ValueError(component)
    config = sdnq_config(skip, device)
    print(f'Quantizing {component}; math on {device}, saved weights on CPU', flush=True)
    original_quantize = _progress(component)
    try:
        model = cls.from_pretrained(str(source / component), quantization_config=config,
                                   dtype=torch.bfloat16, device_map={'': 'cpu'},
                                   local_files_only=True)
    except Exception:
        if out.exists():
            shutil.rmtree(out)
        raise
    finally:
        import sdnq.quantizer as quantizer
        quantizer.sdnq_quantize_layer = original_quantize
    counts = Counter()
    for parameter in model.parameters():
        counts[str(parameter.dtype)] += parameter.numel()
    classes = Counter(type(module).__name__ for module in model.modules())
    quantized = [name for name, module in model.named_modules()
                 if 'SDNQ' in type(module).__name__ or 'Quant' in type(module).__name__]
    if not quantized:
        raise RuntimeError(f'No SDNQ modules produced for {component}')
    out.mkdir(parents=True, exist_ok=True)
    model.save_pretrained(str(out), safe_serialization=True, max_shard_size='5GB')
    saved = json.loads((out / 'config.json').read_text(encoding='utf-8'))
    quant = saved.get('quantization_config') or {}
    if str(quant.get('quant_method', '')).lower() != 'sdnq' or str(quant.get('weights_dtype', '')).lower() != 'uint4':
        raise RuntimeError(f'Saved config is not SDNQ UINT4: {quant}')
    if quant.get('use_quantized_matmul'):
        raise RuntimeError('Refusing to save a CUDA-only quantized-matmul checkpoint')
    report = {'component': component, 'method': 'sdnq', 'weights_dtype': 'uint4',
              'use_svd': True, 'svd_rank': 32, 'use_quantized_matmul': False,
              'skip_modules': skip, 'quantized_modules': quantized,
              'module_types': dict(classes), 'parameter_dtypes': dict(counts),
              'saved_quantization_config': quant}
    (target / f'{component}-quantization.json').write_text(json.dumps(report, indent=2), encoding='utf-8')
    del model
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
    print(f'Saved {component}: {len(quantized)} SDNQ modules', flush=True)
    return report


def copy_support(source, target):
    source, target = Path(source), Path(target)
    target.mkdir(parents=True, exist_ok=True)
    for name in ('processor', 'scheduler', 'vae'):
        shutil.copytree(source / name, target / name, dirs_exist_ok=True,
                        ignore=shutil.ignore_patterns('.cache', '*.lock', '*.incomplete'))
    for name in ('model_index.json', 'LICENSE'):
        shutil.copy2(source / name, target / name)
    (target / 'Notice').write_text(NOTICE, encoding='utf-8')


def verify_source(source, manifest):
    source, manifest = Path(source), Path(manifest)
    payload = json.loads(manifest.read_text(encoding='utf-8'))
    revision = payload.get('revision') or payload.get('sha')
    if revision != REVISION:
        raise ValueError(f'Source manifest revision {revision} != {REVISION}')
    rows = payload.get('files')
    if rows is None:
        rows = [{'path': item['rfilename'], 'size': item['size'], 'sha256': item['lfs']['sha256']}
                for item in payload['siblings'] if item['rfilename'].endswith('.safetensors')]
    if len(rows) != 7:
        raise ValueError('Expected seven official weight files')
    verify_files(source, rows)
    return payload


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--source', default=str(Path('..') / 'Qwen-Image-2.1' / 'models' / 'bf16'))
    parser.add_argument('--output', default='models/int4')
    parser.add_argument('--manifest', default=str(Path('..') / 'Qwen-Image-2.1' / 'artifacts' / 'download-verification.json'))
    parser.add_argument('--component', choices=['transformer', 'text_encoder', 'all'], default='all')
    parser.add_argument('--device', default='cuda:0')
    args = parser.parse_args()
    verify_source(args.source, args.manifest)
    print(f'Source weights match {REVISION}', flush=True)
    copy_support(args.source, args.output)
    components = ['transformer', 'text_encoder'] if args.component == 'all' else [args.component]
    for component in components:
        quantize_component(args.source, args.output, component, args.device)
    print(f'Quantized pipeline written to {args.output}', flush=True)
    print(f'Diffusers commit used by this project: {DIFFUSERS_COMMIT}', flush=True)


if __name__ == '__main__':
    main()