mod_config.py Source File

mod_config.py Source File#

Mobilint SDK qb Compiler: mod_config.py Source File
Mobilint SDK qb Compiler v0.11.0.1
MCS002-EN
mod_config.py
1import numbers
2
3from typeguard import typechecked
4from . import ConfigABC
5
6"""
7MOD_apply: Run Min output difference (MOD) QAT if True
8MOD_epochs: Epochs for MOD QAT
9MOD_warmupEpochs: Warmup epochs for MOD QAT
10MOD_lrMinRatio: Minimum lr ratio for cosine scheduler in MOD QAT
11MOD_actScaleLR: LR of activation scale for MOD QAT
12MOD_zeropointLR: LR of zeropoint for MOD QAT
13MOD_weightScaleLR: LR of weight scale for MOD QAT
14MOD_weightLR: LR of weight for MOD QAT
15MOD_biasLR: LR of bias for MOD QAT
16MOD_batchSize: Batch size for MOD QAT
17MOD_qDrop: QDrop ratio for MOD QAT. range=[0, 1]
18MOD_quantizeWeight: Use quantized weight for MOD QAT if True
19MOD_weightScaleInit: Weight quantization init method for for MOD QAT. Options=('MinMax')
20MOD_downresolMode: Down-resolution of fake-quantization, Options=('STE', 'NIPQ')
21MOD_post: Post-processing function to make the output values, Options=(AnchorlessDetection', 'AnchorlessDetectionNMS', 'AnchorlessDetectionNMSCIoU')
22MOD_boxConfThres: Confidence treshold of NMS if NMS is included in the post-processing.
23MOD_boxIoUThres: IoU treshold of NMS if NMS is included in the post-processing.
24MOD_lossType: Loss funciton to compute the loss of the final output layers. Options=('CE', 'CELabel', 'KL', 'AnchorlessDetection')
25MOD_useOutputs: Compare only intermediate layers if True.
26MOD_KLTemperature: Temperature of KL diveregence loss if it is used.
27MOD_reconProb: Sampling probability of intermediate layers. range=[0, 1]
28MOD_reconCoeff: Coefficient to the loss of intermediate layers.
29MOD_lambda_0: 1st coeff of AnchorlessDetection loss
30MOD_lambda_1: 2nd coeff of AnchorlessDetection loss
31MOD_lambda_2: 3rd coeff of AnchorlessDetection loss
32MOD_seed: Random seed for MOD QAT
33"""
34
35
36class ModConfig(ConfigABC):
37 py_to_cpp_name = {
38 "MOD_apply": "apply",
39 "MOD_epochs": "epochs",
40 "MOD_warmupEpochs": "warmupEpochs",
41 "MOD_lrMinRatio": "lrMinRatio",
42 "MOD_actScaleLR": "actScaleLR",
43 "MOD_zeropointLR": "zeropointLR",
44 "MOD_weightScaleLR": "weightScaleLR",
45 "MOD_weightLR": "weightLR",
46 "MOD_biasLR": "biasLR",
47 "MOD_batchSize": "batchSize",
48 "MOD_qDrop": "qDrop",
49 "MOD_quantizeWeight": "quantizeWeight",
50 "MOD_weightScaleInit": "weightScaleInit",
51 "MOD_downresolMode": "downresolMode",
52 "MOD_post": "post",
53 "MOD_boxConfThres": "boxConfThres",
54 "MOD_boxIoUThres": "boxIoUThres",
55 "MOD_lossType": "lossType",
56 "MOD_useOutputs": "useOutputs",
57 "MOD_KLTemperature": "KLTemperature",
58 "MOD_reconProb": "reconProb",
59 "MOD_reconCoeff": "reconCoeff",
60 "MOD_lambda_0": "lambda_0",
61 "MOD_lambda_1": "lambda_1",
62 "MOD_lambda_2": "lambda_2",
63 "MOD_seed": "seed",
64 }
65 min_max_check = {
66 "MOD_epochs": [0, None],
67 "MOD_warmupEpochs": [0, None],
68 "MOD_lrMinRatio": [0, None],
69 "MOD_actScaleLR": [0, None],
70 "MOD_zeropointLR": [0, None],
71 "MOD_weightScaleLR": [0, None],
72 "MOD_weightLR": [0, None],
73 "MOD_biasLR": [0, None],
74 "MOD_batchSize": [0, None],
75 "MOD_boxConfThres": [0, None],
76 "MOD_boxIoUThres": [0, None],
77 "MOD_KLTemperature": [0, None],
78 "MOD_reconProb": [0, None],
79 "MOD_reconCoeff": [0, None],
80 "MOD_lambda_0": [0, None],
81 "MOD_lambda_1": [0, None],
82 "MOD_lambda_2": [0, None],
83 "MOD_seed": [0, None],
84 }
85
86 @typechecked
87 def __init__(
88 self,
89 MOD_apply: bool = False,
90 MOD_epochs: int = 4,
91 MOD_warmupEpochs: int = 1,
92 MOD_lrMinRatio: numbers.Number = 0.1,
93 MOD_actScaleLR: numbers.Number = 0.0,
94 MOD_zeropointLR: numbers.Number = 0.0,
95 MOD_weightScaleLR: numbers.Number = 0.0,
96 MOD_weightLR: numbers.Number = 4e-6,
97 MOD_biasLR: numbers.Number = 4e-6,
98 MOD_batchSize: int = 1,
99 MOD_qDrop: numbers.Number = 0.0,
100 MOD_quantizeWeight: bool = True,
101 MOD_weightScaleInit: str = "MinMax",
102 MOD_downresolMode: str = "STE",
103 MOD_post: str = "",
104 MOD_boxConfThres: numbers.Number = 0,
105 MOD_boxIoUThres: numbers.Number = 0,
106 MOD_lossType: str = "mse",
107 MOD_useOutputs: bool = False,
108 MOD_KLTemperature: numbers.Number = 1.0,
109 MOD_reconProb: numbers.Number = 1.0,
110 MOD_reconCoeff: numbers.Number = 1.0,
111 MOD_lambda_0: numbers.Number = 1.0,
112 MOD_lambda_1: numbers.Number = 1.0,
113 MOD_lambda_2: numbers.Number = 1.0,
114 MOD_seed: int = 0,
115 ): # Mod is optional
116 super().__init__(optional=True)
117 self.MOD_apply = MOD_apply
118 self.MOD_epochs = MOD_epochs
119 self.MOD_warmupEpochs = MOD_warmupEpochs
120 self.MOD_lrMinRatio = MOD_lrMinRatio
121 self.MOD_actScaleLR = MOD_actScaleLR
122 self.MOD_zeropointLR = MOD_zeropointLR
123 self.MOD_weightScaleLR = MOD_weightScaleLR
124 self.MOD_weightLR = MOD_weightLR
125 self.MOD_biasLR = MOD_biasLR
126 self.MOD_batchSize = MOD_batchSize
127 self.MOD_qDrop = MOD_qDrop
128 self.MOD_quantizeWeight = MOD_quantizeWeight
129 self.MOD_weightScaleInit = MOD_weightScaleInit
130 self.MOD_downresolMode = MOD_downresolMode
131 self.MOD_post = MOD_post
132 self.MOD_boxConfThres = MOD_boxConfThres
133 self.MOD_boxIoUThres = MOD_boxIoUThres
134 self.MOD_lossType = MOD_lossType
135 self.MOD_useOutputs = MOD_useOutputs
136 self.MOD_KLTemperature = MOD_KLTemperature
137 self.MOD_reconProb = MOD_reconProb
138 self.MOD_reconCoeff = MOD_reconCoeff
139 self.MOD_lambda_0 = MOD_lambda_0
140 self.MOD_lambda_1 = MOD_lambda_1
141 self.MOD_lambda_2 = MOD_lambda_2
142 self.MOD_seed = MOD_seed
143 self.check_valid()
144
145 @classmethod
146 def default_config(cls):
147 return ModConfig()
148
149 def to_dict(self):
150 mod_config_dict = dict()
151 mod_params = dict()
152 for py_name, cpp_name in ModConfig.py_to_cpp_name.items():
153 value = getattr(self, py_name)
154 if isinstance(value, str) and cpp_name not in [
155 "weightScaleInit",
156 "downresolMode",
157 ]: # these options are not use lower string in cpp
158 value = value.lower()
159 mod_params[cpp_name] = value
160 mod_config_dict["minOutputDiff"] = mod_params
161
162 return mod_config_dict