FLUX.1-schnell-OpenVINO-INT4 / quantize_int4_flux.py
HelloSun's picture
Add quantize_int4_flux.py
62c7051 verified
Raw History Blame Contribute Delete
2.79 kB
"""NNCF weight-only INT4 quantization of the OpenVINO FP16 FLUX pipeline.
transformer + text_encoder -> weight-only INT4 (asymmetric, group 128)
all remaining components -> default weight-only INT8
Usage:
python quantize_int4_flux.py \
--model_path /home/user/app/flux-schnell-ov-fp16 \
--output_path /home/user/app/flux-schnell-ov-int4
"""
import argparse
import shutil
import time
from pathlib import Path
from optimum.intel import OVConfig, OVQuantizer
from optimum.intel.openvino import (
OVPipelineQuantizationConfig,
OVWeightQuantizationConfig,
)
# For FLUX we need to import the specific pipeline class
try:
from optimum.intel import OVFluxPipeline
except ImportError:
# Fallback to generic diffusion pipeline
from optimum.intel import OVDiffusionPipeline as OVFluxPipeline
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_path",
type=str,
default="/home/user/app/flux-schnell-ov-fp16",
help="OpenVINO FP16 pipeline exported with optimum-cli",
)
parser.add_argument(
"--output_path",
type=str,
default="/home/user/app/flux-schnell-ov-int4",
help="Destination folder for the INT4 pipeline",
)
args = parser.parse_args()
model_path = Path(args.model_path)
output_path = Path(args.output_path)
if output_path.exists():
shutil.rmtree(output_path)
output_path.mkdir(parents=True, exist_ok=True)
int4 = dict(
bits=4,
sym=False,
group_size=128,
group_size_fallback="adjust",
ratio=1.0,
)
quantization_configs = {
# FLUX uses "transformer" instead of "unet"
"transformer": OVWeightQuantizationConfig(**int4),
"text_encoder": OVWeightQuantizationConfig(**int4),
# Also quantize the second text encoder (T5)
"text_encoder_2": OVWeightQuantizationConfig(**int4),
}
default_config = OVWeightQuantizationConfig(bits=8)
quantization_config = OVPipelineQuantizationConfig(
quantization_configs=quantization_configs,
default_config=default_config,
)
ov_config = OVConfig(quantization_config=quantization_config)
print(f"loading FP16 pipeline from {model_path} ...", flush=True)
t0 = time.perf_counter()
model = OVFluxPipeline.from_pretrained(str(model_path), device="CPU")
print(f"loaded in {time.perf_counter() - t0:.1f}s", flush=True)
quantizer = OVQuantizer(model=model)
t0 = time.perf_counter()
quantizer.quantize(ov_config=ov_config, save_directory=str(output_path))
print(f"quantization took {time.perf_counter() - t0:.1f}s", flush=True)
print(f"INT4 pipeline saved to {output_path}")
if __name__ == "__main__":
main()