HelloSun commited on
Commit
62c7051
·
verified ·
1 Parent(s): 7fd6408

Add quantize_int4_flux.py

Browse files
Files changed (1) hide show
  1. quantize_int4_flux.py +88 -0
quantize_int4_flux.py ADDED
@@ -0,0 +1,88 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """NNCF weight-only INT4 quantization of the OpenVINO FP16 FLUX pipeline.
2
+
3
+ transformer + text_encoder -> weight-only INT4 (asymmetric, group 128)
4
+ all remaining components -> default weight-only INT8
5
+
6
+ Usage:
7
+ python quantize_int4_flux.py \
8
+ --model_path /home/user/app/flux-schnell-ov-fp16 \
9
+ --output_path /home/user/app/flux-schnell-ov-int4
10
+ """
11
+
12
+ import argparse
13
+ import shutil
14
+ import time
15
+ from pathlib import Path
16
+
17
+ from optimum.intel import OVConfig, OVQuantizer
18
+ from optimum.intel.openvino import (
19
+ OVPipelineQuantizationConfig,
20
+ OVWeightQuantizationConfig,
21
+ )
22
+
23
+ # For FLUX we need to import the specific pipeline class
24
+ try:
25
+ from optimum.intel import OVFluxPipeline
26
+ except ImportError:
27
+ # Fallback to generic diffusion pipeline
28
+ from optimum.intel import OVDiffusionPipeline as OVFluxPipeline
29
+
30
+
31
+ def main() -> None:
32
+ parser = argparse.ArgumentParser()
33
+ parser.add_argument(
34
+ "--model_path",
35
+ type=str,
36
+ default="/home/user/app/flux-schnell-ov-fp16",
37
+ help="OpenVINO FP16 pipeline exported with optimum-cli",
38
+ )
39
+ parser.add_argument(
40
+ "--output_path",
41
+ type=str,
42
+ default="/home/user/app/flux-schnell-ov-int4",
43
+ help="Destination folder for the INT4 pipeline",
44
+ )
45
+ args = parser.parse_args()
46
+
47
+ model_path = Path(args.model_path)
48
+ output_path = Path(args.output_path)
49
+ if output_path.exists():
50
+ shutil.rmtree(output_path)
51
+ output_path.mkdir(parents=True, exist_ok=True)
52
+
53
+ int4 = dict(
54
+ bits=4,
55
+ sym=False,
56
+ group_size=128,
57
+ group_size_fallback="adjust",
58
+ ratio=1.0,
59
+ )
60
+ quantization_configs = {
61
+ # FLUX uses "transformer" instead of "unet"
62
+ "transformer": OVWeightQuantizationConfig(**int4),
63
+ "text_encoder": OVWeightQuantizationConfig(**int4),
64
+ # Also quantize the second text encoder (T5)
65
+ "text_encoder_2": OVWeightQuantizationConfig(**int4),
66
+ }
67
+ default_config = OVWeightQuantizationConfig(bits=8)
68
+
69
+ quantization_config = OVPipelineQuantizationConfig(
70
+ quantization_configs=quantization_configs,
71
+ default_config=default_config,
72
+ )
73
+ ov_config = OVConfig(quantization_config=quantization_config)
74
+
75
+ print(f"loading FP16 pipeline from {model_path} ...", flush=True)
76
+ t0 = time.perf_counter()
77
+ model = OVFluxPipeline.from_pretrained(str(model_path), device="CPU")
78
+ print(f"loaded in {time.perf_counter() - t0:.1f}s", flush=True)
79
+
80
+ quantizer = OVQuantizer(model=model)
81
+ t0 = time.perf_counter()
82
+ quantizer.quantize(ov_config=ov_config, save_directory=str(output_path))
83
+ print(f"quantization took {time.perf_counter() - t0:.1f}s", flush=True)
84
+ print(f"INT4 pipeline saved to {output_path}")
85
+
86
+
87
+ if __name__ == "__main__":
88
+ main()