Image21-INT8 / scripts /quantize.py
ixim's picture
Fix INT8 CPU offload; publish memory audit and paired retests
ba48d54 verified
Raw
History Blame Contribute Delete
4.21 kB
"""Serialize standard bitsandbytes INT8 components for a Diffusers pipeline."""
import argparse
import gc
import json
import shutil
from collections import Counter
from pathlib import Path
import torch
from diffusers import BitsAndBytesConfig as DBitsAndBytesConfig
from diffusers import QwenImage21Transformer2DModel
from transformers import BitsAndBytesConfig as TBitsAndBytesConfig
from transformers import Qwen3VLForConditionalGeneration
TRANSFORMER_SKIP = ['proj_out', 'img_in', 'txt_in', 'time_text_embed', 'modulation', 'norm_out']
TEXT_SKIP = ['lm_head', 'visual']
def quantize_component(source, target, component):
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':
cls, config_cls, skip = QwenImage21Transformer2DModel, DBitsAndBytesConfig, TRANSFORMER_SKIP
elif component == 'text_encoder':
cls, config_cls, skip = Qwen3VLForConditionalGeneration, TBitsAndBytesConfig, TEXT_SKIP
else:
raise ValueError(component)
config = config_cls(load_in_8bit=True, llm_int8_threshold=6.0,
llm_int8_skip_modules=skip)
model = cls.from_pretrained(str(source / component), quantization_config=config,
dtype=torch.bfloat16, device_map={'': 0},
local_files_only=True)
import bitsandbytes as bnb
modules = [name for name, mod in model.named_modules() if isinstance(mod, bnb.nn.Linear8bitLt)]
counts = Counter()
for p in model.parameters():
counts[str(p.dtype)] += p.numel()
report = {'component': component, 'method': 'bitsandbytes LLM.int8',
'threshold': 6.0, 'skip_modules': skip, 'linear8bit_modules': modules,
'parameter_dtypes': dict(counts), 'int8_parameters': counts['torch.int8']}
if not modules or not report['int8_parameters']:
raise RuntimeError(f'No integer weights produced for {component}')
out.mkdir(parents=True, exist_ok=True)
model.save_pretrained(str(out), safe_serialization=True, max_shard_size='5GB')
(target / f'{component}-quantization.json').write_text(
json.dumps(report, indent=2), encoding='utf-8')
del model
gc.collect()
torch.cuda.empty_cache()
return report
def copy_support(source, 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(
'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-INT8, by ixim / iximbox.\n'
'Modified files: transformer and text_encoder weight shards, shard indexes and '
'config.json files. Eligible linear weights converted to bitsandbytes INT8; '
'floating-point exceptions are documented in component quantization reports.\n'
'VAE, scheduler and processor retained from the pinned upstream snapshot.\n',
encoding='utf-8')
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--source', default='models/bf16')
ap.add_argument('--output', default='models/int8')
ap.add_argument('--component', choices=['transformer', 'text_encoder', 'all'], default='all')
args = ap.parse_args()
components = ['transformer', 'text_encoder'] if args.component == 'all' else [args.component]
for component in components:
print(f'Quantizing {component}', flush=True)
report = quantize_component(args.source, args.output, component)
print(f'Saved {component}: {report["int8_parameters"]:,} INT8 parameters', flush=True)
copy_support(Path(args.source), Path(args.output))
if __name__ == '__main__':
main()