"""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()