Image21-INT4 / scripts /quantize.py
ixim's picture
Release verified Image21-INT4 conversion
9116984 verified
Raw
History Blame Contribute Delete
7.76 kB
"""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()