File size: 5,291 Bytes
77a7159
 
 
 
 
 
 
 
 
 
 
15b2edb
77a7159
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b42379b
 
 
 
 
 
 
 
 
 
77a7159
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
15b2edb
c4942a1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77a7159
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
import os
import glob
import sys
import argparse
import cv2
import numpy as np
import onnx
from onnxruntime.quantization import quantize_static, CalibrationDataReader, QuantType

class CodeFormerCalibrationDataReader(CalibrationDataReader):
    """Calibration Data Reader for ONNX Runtime static quantization of CodeFormer."""
    def __init__(self, calibration_folder: str, w_val: float = 0.5, max_samples: int = 5):
        super().__init__()
        valid_exts = ('.png', '.jpg', '.jpeg', '.webp')
        self.image_paths = [
            os.path.join(calibration_folder, f) for f in os.listdir(calibration_folder)
            if f.lower().endswith(valid_exts)
        ] if os.path.exists(calibration_folder) else []

        self.w_val = np.array([w_val], dtype=np.float32)
        self.enum_data_dicts = []
        self._preprocess(max_samples)

    def _preprocess(self, max_samples: int):
        print(f"[Calib] Preprocessing up to {max_samples} calibration images...")
        count = 0
        for img_path in self.image_paths:
            if count >= max_samples:
                break
            img = cv2.imread(img_path)
            if img is None:
                continue
            img = cv2.resize(img, (512, 512), interpolation=cv2.INTER_LANCZOS4)
            img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0
            # Normalize to [-1, 1] range matching CodeFormer inputs
            img_norm = (img_rgb - 0.5) / 0.5
            tensor = np.transpose(img_norm, (2, 0, 1))[np.newaxis, ...].astype(np.float32)
            self.enum_data_dicts.append({
                "input": tensor,
                "w": self.w_val
            })
            count += 1

        if len(self.enum_data_dicts) == 0:
            print("[Calib] No calibration images found. Generating 10 synthetic calibration samples...")
            np.random.seed(42)
            for _ in range(10):
                dummy_tensor = np.random.uniform(-1.0, 1.0, (1, 3, 512, 512)).astype(np.float32)
                self.enum_data_dicts.append({
                    "input": dummy_tensor,
                    "w": self.w_val
                })

        print(f"[Calib] Prepared {len(self.enum_data_dicts)} calibration samples.")
        self.enum_data = iter(self.enum_data_dicts)

    def get_next(self):
        return next(self.enum_data, None)

    def rewind(self):
        self.enum_data = iter(self.enum_data_dicts)

def quantize_static_model(model_path: str, output_path: str, calibration_folder: str):
    if not os.path.isfile(model_path):
        raise FileNotFoundError(f"Model file not found: {model_path}")

    dr = CodeFormerCalibrationDataReader(calibration_folder)
    if dr.datasize == 0 if hasattr(dr, 'datasize') else len(dr.enum_data_dicts) == 0:
        print(f"[ERROR] Calibration folder is empty or invalid: {calibration_folder}")
        sys.exit(1)

    print(f"[Static Quant] Quantizing {model_path} -> {output_path}")
    try:
        quantize_static(
            model_input=model_path,
            model_output=output_path,
            calibration_data_reader=dr,
            quant_format=QuantType.QInt8,
            activation_type=QuantType.QUInt8,
            weight_type=QuantType.QInt8,
            use_external_data_format=True
        )
        # Verify the quantized model can be loaded cleanly by ONNX Runtime
        import onnxruntime as ort
        opts = ort.SessionOptions()
        opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
        test_session = ort.InferenceSession(output_path, sess_options=opts, providers=['CPUExecutionProvider'])
        print(f"[Static Quant] Verification passed! Model loaded successfully: {output_path}")
    except Exception as e:
        print(f"[Static Quant] Static quantization verification failed: {e}")
        print(f"[Static Quant] Falling back to dynamic INT8 quantization...")
        from onnxruntime.quantization import quantize_dynamic
        quantize_dynamic(
            model_input=model_path,
            model_output=output_path,
            weight_type=QuantType.QInt8,
            use_external_data_format=True
        )
        print(f"[Static Quant] Dynamic INT8 quantization fallback completed: {output_path}")

if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="Static INT8 quantization for CodeFormer ONNX with calibration.")
    parser.add_argument("--model", type=str, default="weights/CodeFormer/codeformer.onnx", help="Path to input ONNX model.")
    parser.add_argument("--output", type=str, default="weights/CodeFormer/codeformer_int8_v2.onnx", help="Path to save output static quantized model.")
    parser.add_argument("--calib-dir", type=str, default="models/CodeFormer/datasets/ffhq/ffhq_512", help="Calibration images directory.")
    args = parser.parse_args()

    project_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
    model_p = os.path.join(project_dir, args.model) if not os.path.isabs(args.model) else args.model
    out_p = os.path.join(project_dir, args.output) if not os.path.isabs(args.output) else args.output
    calib_p = os.path.join(project_dir, args.calib_dir) if not os.path.isabs(args.calib_dir) else args.calib_dir

    quantize_static_model(model_p, out_p, calib_p)