equivalent_transformation_config.py Source File

equivalent_transformation_config.py Source File#

Mobilint SDK qb Compiler: equivalent_transformation_config.py Source File
Mobilint SDK qb Compiler v0.12.0.0
MCS002-KR
equivalent_transformation_config.py
1from typing import Dict, List, Optional
2from typeguard import typechecked
3from . import ConfigABC
4from .config_utils import value_range_check
5
6
12
13
15 """
16 @brief Configuration for equivalent transformation techniques.
17
18 @details This includes various transformation methods like NormConv, QK, UD, VO, LUTRescale,
19 FeedForwardMultiLUT, SpinR1, HeadOutChRotation, SpinR2, QKRotation, and FlattenQuant.
20 """
21
22 optional = False
23 py_to_cpp_name = {}
24 # Empty py_to_cpp_name since this config uses nested structure in to_dict
25
26 # Default values shared between __init__ and from_dict
27 DEFAULTS = {
28 "seed": 0,
29 "layer_name_pattern": {"head": ""},
30 "apply_hadamard_rotation_matrix": True,
31 # NormConv
32 "norm_conv_apply": False,
33 "norm_conv_learn": False,
34 "norm_conv_smoothing_factor": 0.5,
35 "norm_conv_min_gamma": 0.0001,
36 "norm_conv_max_gamma": 10000.0,
37 # QK
38 "qk_apply": False,
39 "qk_smoothing_factor": 0.5,
40 "qk_min_gamma": 0.0001,
41 "qk_max_gamma": 10000.0,
42 # UD
43 "ud_apply": False,
44 "ud_learn": False,
45 "ud_smoothing_factor": 0.5,
46 "ud_min_gamma": 0.0001,
47 "ud_max_gamma": 10000.0,
48 # VO
49 "vo_apply": False,
50 "vo_smoothing_factor": 0.5,
51 "vo_min_gamma": 0.0001,
52 "vo_max_gamma": 10000.0,
53 # LUTRescale
54 "lut_rescale_apply": False,
55 "lut_rescale_alpha": True,
56 "lut_rescale_beta": False,
57 # FeedForwardMultiLUT
58 "ff_multi_lut_apply": False,
59 "ff_multi_lut_breakpoints": [-8.0, -4.0, 0],
60 # SpinR1
61 "spin_r1_apply": False,
62 "spin_r1_matrix_path": "",
63 # HeadOutChRotation
64 "head_out_ch_rotation_apply": False,
65 "head_out_ch_rotation_matrix_path": "",
66 # SpinR2
67 "spin_r2_apply": False,
68 "spin_r2_learn": False,
69 "spin_r2_matrix_path": "",
70 # QKRotation
71 "qk_rotation_apply": False,
72 "qk_rotation_matrix_path": "",
73 # FlattenQuant
74 "flatten_quant_apply": False,
75 "flatten_quant_learn": False,
76 "flatten_quant_apply_threshold": 0.33,
77 "flatten_quant_max_overhead": 0.05,
78 # SplitFFN
79 "split_ffn_apply": False,
80 "split_ffn_ch_per_ffn": -1,
81 }
82
83 # Min/max range checks
84 min_max_check = {
85 "norm_conv_smoothing_factor": [0.0, 1.0],
86 "norm_conv_min_gamma": [0.0, None],
87 "norm_conv_max_gamma": [0.0, None],
88 "qk_smoothing_factor": [0.0, 1.0],
89 "qk_min_gamma": [0.0, None],
90 "qk_max_gamma": [0.0, None],
91 "ud_smoothing_factor": [0.0, 1.0],
92 "ud_min_gamma": [0.0, None],
93 "ud_max_gamma": [0.0, None],
94 "vo_smoothing_factor": [0.0, 1.0],
95 "vo_min_gamma": [0.0, None],
96 "vo_max_gamma": [0.0, None],
97 "flatten_quant_apply_threshold": [0.0, 1.0],
98 "flatten_quant_max_overhead": [0.0, None],
99 }
100
102 self,
103 seed: int,
104 layer_name_pattern: Dict[str, str],
105 apply_hadamard_rotation_matrix: bool,
106 # NormConv
107 norm_conv_apply: bool,
108 norm_conv_learn: bool,
109 norm_conv_smoothing_factor: float,
110 norm_conv_min_gamma: float,
111 norm_conv_max_gamma: float,
112 # QK
113 qk_apply: bool,
114 qk_smoothing_factor: float,
115 qk_min_gamma: float,
116 qk_max_gamma: float,
117 # UD
118 ud_apply: bool,
119 ud_learn: bool,
120 ud_smoothing_factor: float,
121 ud_min_gamma: float,
122 ud_max_gamma: float,
123 # VO
124 vo_apply: bool,
125 vo_smoothing_factor: float,
126 vo_min_gamma: float,
127 vo_max_gamma: float,
128 # LUTRescale
129 lut_rescale_apply: bool,
130 lut_rescale_alpha: bool,
131 lut_rescale_beta: bool,
132 # FeedForwardMultiLUT
133 ff_multi_lut_apply: bool,
134 ff_multi_lut_breakpoints: List[float],
135 # SpinR1
136 spin_r1_apply: bool,
137 spin_r1_matrix_path: str,
138 # HeadOutChRotation
139 head_out_ch_rotation_apply: bool,
140 head_out_ch_rotation_matrix_path: str,
141 # SpinR2
142 spin_r2_apply: bool,
143 spin_r2_learn: bool,
144 spin_r2_matrix_path: str,
145 # QKRotation
146 qk_rotation_apply: bool,
147 qk_rotation_matrix_path: str,
148 # FlattenQuant
149 flatten_quant_apply: bool,
150 flatten_quant_learn: bool,
151 flatten_quant_apply_threshold: float,
152 flatten_quant_max_overhead: float,
153 # SplitFFN
154 split_ffn_apply: bool,
155 split_ffn_ch_per_ffn: int,
156 optional: bool = optional,
157 ):
158 """
159 @brief Initialize the EquivalentTransformationConfig.
160 @param seed int. Random seed used for equivalent transformation operations. Default is 0.
161 @param layer_name_pattern Dict[str, str]. Mapping from logical layer groups (e.g., "head") to name patterns used to match target layers. Default is {"head": ""}.
162 @param apply_hadamard_rotation_matrix bool. If True, apply the Hadamard rotation matrix where applicable. Default is True.
163
164 # NormConv
165 @param norm_conv_apply bool. If True, enable NormConv transformation. Default is False.
166 @param norm_conv_learn bool. If True, learn NormConv parameters during calibration/optimization. Default is False.
167 @param norm_conv_smoothing_factor float. Smoothing factor used in NormConv. Default is 0.5.
168 @param norm_conv_min_gamma float. Minimum gamma value (lower bound clamp) for NormConv. Default is 0.0001.
169 @param norm_conv_max_gamma float. Maximum gamma value (upper bound clamp) for NormConv. Default is 10000.0.
170
171 # QK
172 @param qk_apply bool. If True, enable QK transformation. Default is False.
173 @param qk_smoothing_factor float. Smoothing factor used in QK. Default is 0.5.
174 @param qk_min_gamma float. Minimum gamma value (lower bound clamp) for QK. Default is 0.0001.
175 @param qk_max_gamma float. Maximum gamma value (upper bound clamp) for QK. Default is 10000.0.
176
177 # UD
178 @param ud_apply bool. If True, enable UD transformation. Default is False.
179 @param ud_learn bool. If True, learn UD parameters during calibration/optimization. Default is False.
180 @param ud_smoothing_factor float. Smoothing factor used in UD. Default is 0.5.
181 @param ud_min_gamma float. Minimum gamma value (lower bound clamp) for UD. Default is 0.0001.
182 @param ud_max_gamma float. Maximum gamma value (upper bound clamp) for UD. Default is 10000.0.
183
184 # VO
185 @param vo_apply bool. If True, enable VO transformation. Default is False.
186 @param vo_smoothing_factor float. Smoothing factor used in VO. Default is 0.5.
187 @param vo_min_gamma float. Minimum gamma value (lower bound clamp) for VO. Default is 0.0001.
188 @param vo_max_gamma float. Maximum gamma value (upper bound clamp) for VO. Default is 10000.0.
189
190 # LUTRescale
191 @param lut_rescale_apply bool. If True, enable LUT rescaling. Default is False.
192 @param lut_rescale_alpha bool. If True, apply alpha rescaling in LUTRescale. Default is True.
193 @param lut_rescale_beta bool. If True, apply beta rescaling in LUTRescale. Default is False.
194
195 # FeedForwardMultiLUT
196 @param ff_multi_lut_apply bool. If True, enable FeedForward Multi-LUT. Default is False.
197 @param ff_multi_lut_breakpoints List[float]. Breakpoints used for FeedForward Multi-LUT piecewise regions. Default is [-8.0, -4.0, 0].
198
199 # SpinR1
200 @param spin_r1_apply bool. If True, enable SpinR1 rotation. Default is False.
201 @param spin_r1_matrix_path str. Path to the SpinR1 rotation matrix file. Default is "".
202
203 # HeadOutChRotation
204 @param head_out_ch_rotation_apply bool. If True, enable HeadOutChRotation. Default is False.
205 @param head_out_ch_rotation_matrix_path str. Path to the HeadOutChRotation matrix file. Default is "".
206
207 # SpinR2
208 @param spin_r2_apply bool. If True, enable SpinR2 rotation. Default is False.
209 @param spin_r2_learn bool. If True, learn the SpinR2 rotation matrix/parameters during calibration/optimization. Default is False.
210 @param spin_r2_matrix_path str. Path to the SpinR2 rotation matrix file. Default is "".
211
212 # QKRotation
213 @param qk_rotation_apply bool. If True, enable QKRotation. Default is False.
214 @param qk_rotation_matrix_path str. Path to the QKRotation matrix file. Default is "".
215
216 # FlattenQuant
217 @param flatten_quant_apply bool. If True, enable FlattenQuant. Default is False.
218 @param flatten_quant_learn bool. If True, learn FlattenQuant parameters during calibration/optimization. Default is False.
219 @param flatten_quant_apply_threshold float. Threshold for applying FlattenQuant. Default is 0.33.
220 @param flatten_quant_max_overhead float. Maximum allowed overhead when applying FlattenQuant. Default is 0.05.
221
222 # SplitFFN
223 @param split_ffn_apply bool. If True, enable SplitFFN. Default is False.
224 @param split_ffn_ch_per_ffn int. Number of channels per FFN used when SplitFFN is enabled. Default is -1.
225
226 @param optional bool. Indicates whether this config is optional. Use the default value defined as a class variable.
227 """
228 super().__init__(optional)
229
230 self.seed = seed
231 self.layer_name_pattern = layer_name_pattern
232 self.apply_hadamard_rotation_matrix = apply_hadamard_rotation_matrix
233
234 # NormConv
235 self.norm_conv_apply = norm_conv_apply
236 self.norm_conv_learn = norm_conv_learn
237 self.norm_conv_smoothing_factor = norm_conv_smoothing_factor
238 self.norm_conv_min_gamma = norm_conv_min_gamma
239 self.norm_conv_max_gamma = norm_conv_max_gamma
240
241 # QK
242 self.qk_apply = qk_apply
243 self.qk_smoothing_factor = qk_smoothing_factor
244 self.qk_min_gamma = qk_min_gamma
245 self.qk_max_gamma = qk_max_gamma
246
247 # UD
248 self.ud_apply = ud_apply
249 self.ud_learn = ud_learn
250 self.ud_smoothing_factor = ud_smoothing_factor
251 self.ud_min_gamma = ud_min_gamma
252 self.ud_max_gamma = ud_max_gamma
253
254 # VO
255 self.vo_apply = vo_apply
256 self.vo_smoothing_factor = vo_smoothing_factor
257 self.vo_min_gamma = vo_min_gamma
258 self.vo_max_gamma = vo_max_gamma
259
260 # LUTRescale
261 self.lut_rescale_apply = lut_rescale_apply
262 self.lut_rescale_alpha = lut_rescale_alpha
263 self.lut_rescale_beta = lut_rescale_beta
264
265 # FeedForwardMultiLUT
266 self.ff_multi_lut_apply = ff_multi_lut_apply
267 self.ff_multi_lut_breakpoints = ff_multi_lut_breakpoints
268
269 # SpinR1
270 self.spin_r1_apply = spin_r1_apply
271 self.spin_r1_matrix_path = spin_r1_matrix_path
272
273 # HeadOutChRotation
274 self.head_out_ch_rotation_apply = head_out_ch_rotation_apply
275 self.head_out_ch_rotation_matrix_path = head_out_ch_rotation_matrix_path
276
277 # SpinR2
278 self.spin_r2_apply = spin_r2_apply
279 self.spin_r2_learn = spin_r2_learn
280 self.spin_r2_matrix_path = spin_r2_matrix_path
281
282 # QKRotation
283 self.qk_rotation_apply = qk_rotation_apply
284 self.qk_rotation_matrix_path = qk_rotation_matrix_path
285
286 # FlattenQuant
287 self.flatten_quant_apply = flatten_quant_apply
288 self.flatten_quant_learn = flatten_quant_learn
289 self.flatten_quant_apply_threshold = flatten_quant_apply_threshold
290 self.flatten_quant_max_overhead = flatten_quant_max_overhead
291
292 # SplitFFN
293 self.split_ffn_apply = split_ffn_apply
294 self.split_ffn_ch_per_ffn = split_ffn_ch_per_ffn
295
296 self.check_valid()
297
298 def check_valid(self, except_list=["optional"]):
299 """Override to skip py_to_cpp_name check (uses nested structure in to_dict)"""
300 self.check_min_max_range()
301 self.check_valid_str_configs()
302
303 @classmethod
304 def default_config(cls):
305 return cls(**cls.DEFAULTS, optional=cls.optional)
306
307 @classmethod
308 def from_dict(cls, config_dict: dict):
309 """Create EquivalentTransformationConfig from dictionary (JSON structure)"""
310 # Start with defaults and update with provided values
311 params = cls.DEFAULTS.copy()
312
313 # Top-level parameters
314 if "seed" in config_dict:
315 params["seed"] = config_dict["seed"]
316 if "layerNamePattern" in config_dict:
317 params["layer_name_pattern"] = config_dict["layerNamePattern"]
318 if "applyHadamardRotationMatrix" in config_dict:
319 params["apply_hadamard_rotation_matrix"] = config_dict[
320 "applyHadamardRotationMatrix"
321 ]
322
323 # NormConv
324 if "NormConv" in config_dict:
325 norm_conv = config_dict["NormConv"]
326 if "apply" in norm_conv:
327 params["norm_conv_apply"] = norm_conv["apply"]
328 if "learn" in norm_conv:
329 params["norm_conv_learn"] = norm_conv["learn"]
330 if "smoothingFactor" in norm_conv:
331 params["norm_conv_smoothing_factor"] = norm_conv["smoothingFactor"]
332 if "minGamma" in norm_conv:
333 params["norm_conv_min_gamma"] = norm_conv["minGamma"]
334 if "maxGamma" in norm_conv:
335 params["norm_conv_max_gamma"] = norm_conv["maxGamma"]
336
337 # QK
338 if "QK" in config_dict:
339 qk = config_dict["QK"]
340 if "apply" in qk:
341 params["qk_apply"] = qk["apply"]
342 if "smoothingFactor" in qk:
343 params["qk_smoothing_factor"] = qk["smoothingFactor"]
344 if "minGamma" in qk:
345 params["qk_min_gamma"] = qk["minGamma"]
346 if "maxGamma" in qk:
347 params["qk_max_gamma"] = qk["maxGamma"]
348
349 # UD
350 if "UD" in config_dict:
351 ud = config_dict["UD"]
352 if "apply" in ud:
353 params["ud_apply"] = ud["apply"]
354 if "learn" in ud:
355 params["ud_learn"] = ud["learn"]
356 if "smoothingFactor" in ud:
357 params["ud_smoothing_factor"] = ud["smoothingFactor"]
358 if "minGamma" in ud:
359 params["ud_min_gamma"] = ud["minGamma"]
360 if "maxGamma" in ud:
361 params["ud_max_gamma"] = ud["maxGamma"]
362
363 # VO
364 if "VO" in config_dict:
365 vo = config_dict["VO"]
366 if "apply" in vo:
367 params["vo_apply"] = vo["apply"]
368 if "smoothingFactor" in vo:
369 params["vo_smoothing_factor"] = vo["smoothingFactor"]
370 if "minGamma" in vo:
371 params["vo_min_gamma"] = vo["minGamma"]
372 if "maxGamma" in vo:
373 params["vo_max_gamma"] = vo["maxGamma"]
374
375 # LUTRescale
376 if "LUTRescale" in config_dict:
377 lut_rescale = config_dict["LUTRescale"]
378 if "apply" in lut_rescale:
379 params["lut_rescale_apply"] = lut_rescale["apply"]
380 if "rescaleAlpha" in lut_rescale:
381 params["lut_rescale_alpha"] = lut_rescale["rescaleAlpha"]
382 if "rescaleBeta" in lut_rescale:
383 params["lut_rescale_beta"] = lut_rescale["rescaleBeta"]
384
385 # FeedForwardMultiLUT
386 if "FeedForwardMultiLUT" in config_dict:
387 ff_multi_lut = config_dict["FeedForwardMultiLUT"]
388 if "apply" in ff_multi_lut:
389 params["ff_multi_lut_apply"] = ff_multi_lut["apply"]
390 if "breakpoints" in ff_multi_lut:
391 params["ff_multi_lut_breakpoints"] = ff_multi_lut["breakpoints"]
392
393 # SpinR1
394 if "SpinR1" in config_dict:
395 spin_r1 = config_dict["SpinR1"]
396 if "apply" in spin_r1:
397 params["spin_r1_apply"] = spin_r1["apply"]
398 if "matrixPath" in spin_r1:
399 params["spin_r1_matrix_path"] = spin_r1["matrixPath"]
400
401 # HeadOutChRotation
402 if "HeadOutChRotation" in config_dict:
403 head_out_ch = config_dict["HeadOutChRotation"]
404 if "apply" in head_out_ch:
405 params["head_out_ch_rotation_apply"] = head_out_ch["apply"]
406 if "matrixPath" in head_out_ch:
407 params["head_out_ch_rotation_matrix_path"] = head_out_ch["matrixPath"]
408
409 # SpinR2
410 if "SpinR2" in config_dict:
411 spin_r2 = config_dict["SpinR2"]
412 if "apply" in spin_r2:
413 params["spin_r2_apply"] = spin_r2["apply"]
414 if "learn" in spin_r2:
415 params["spin_r2_learn"] = spin_r2["learn"]
416 if "matrixPath" in spin_r2:
417 params["spin_r2_matrix_path"] = spin_r2["matrixPath"]
418
419 # QKRotation
420 if "QKRotation" in config_dict:
421 qk_rotation = config_dict["QKRotation"]
422 if "apply" in qk_rotation:
423 params["qk_rotation_apply"] = qk_rotation["apply"]
424 if "matrixPath" in qk_rotation:
425 params["qk_rotation_matrix_path"] = qk_rotation["matrixPath"]
426
427 # FlattenQuant
428 if "FlattenQuant" in config_dict:
429 flatten_quant = config_dict["FlattenQuant"]
430 if "apply" in flatten_quant:
431 params["flatten_quant_apply"] = flatten_quant["apply"]
432 if "learn" in flatten_quant:
433 params["flatten_quant_learn"] = flatten_quant["learn"]
434 if "applyThreshold" in flatten_quant:
435 params["flatten_quant_apply_threshold"] = flatten_quant[
436 "applyThreshold"
437 ]
438 if "maxOverhead" in flatten_quant:
439 params["flatten_quant_max_overhead"] = flatten_quant["maxOverhead"]
440
441 # SplitFFN
442 if "SplitFFN" in config_dict:
443 split_ffn = config_dict["SplitFFN"]
444 if "apply" in split_ffn:
445 params["split_ffn_apply"] = split_ffn["apply"]
446 if "chPerFFN" in split_ffn:
447 params["split_ffn_ch_per_ffn"] = split_ffn["chPerFFN"]
448
449 return cls(**params)
450
451 def to_dict(self):
452 return {
453 "seed": self.seed,
454 "layerNamePattern": self.layer_name_pattern,
455 "applyHadamardRotationMatrix": self.apply_hadamard_rotation_matrix,
456 "NormConv": {
457 "apply": self.norm_conv_apply,
458 "learn": self.norm_conv_learn,
459 "smoothingFactor": self.norm_conv_smoothing_factor,
460 "minGamma": self.norm_conv_min_gamma,
461 "maxGamma": self.norm_conv_max_gamma,
462 },
463 "QK": {
464 "apply": self.qk_apply,
465 "smoothingFactor": self.qk_smoothing_factor,
466 "minGamma": self.qk_min_gamma,
467 "maxGamma": self.qk_max_gamma,
468 },
469 "UD": {
470 "apply": self.ud_apply,
471 "learn": self.ud_learn,
472 "smoothingFactor": self.ud_smoothing_factor,
473 "minGamma": self.ud_min_gamma,
474 "maxGamma": self.ud_max_gamma,
475 },
476 "VO": {
477 "apply": self.vo_apply,
478 "smoothingFactor": self.vo_smoothing_factor,
479 "minGamma": self.vo_min_gamma,
480 "maxGamma": self.vo_max_gamma,
481 },
482 "LUTRescale": {
483 "apply": self.lut_rescale_apply,
484 "rescaleAlpha": self.lut_rescale_alpha,
485 "rescaleBeta": self.lut_rescale_beta,
486 },
487 "FeedForwardMultiLUT": {
488 "apply": self.ff_multi_lut_apply,
489 "breakpoints": self.ff_multi_lut_breakpoints,
490 },
491 "SpinR1": {
492 "apply": self.spin_r1_apply,
493 "matrixPath": self.spin_r1_matrix_path,
494 },
495 "HeadOutChRotation": {
496 "apply": self.head_out_ch_rotation_apply,
497 "matrixPath": self.head_out_ch_rotation_matrix_path,
498 },
499 "SpinR2": {
500 "apply": self.spin_r2_apply,
501 "learn": self.spin_r2_learn,
502 "matrixPath": self.spin_r2_matrix_path,
503 },
504 "QKRotation": {
505 "apply": self.qk_rotation_apply,
506 "matrixPath": self.qk_rotation_matrix_path,
507 },
508 "FlattenQuant": {
509 "apply": self.flatten_quant_apply,
510 "learn": self.flatten_quant_learn,
511 "applyThreshold": self.flatten_quant_apply_threshold,
512 "maxOverhead": self.flatten_quant_max_overhead,
513 },
514 "SplitFFN": {
515 "apply": self.split_ffn_apply,
516 "chPerFFN": self.split_ffn_ch_per_ffn,
517 },
518 }
519
520
521def get_equivalent_transformation_config(**kwargs) -> EquivalentTransformationConfig:
522 """
523 @brief Create EquivalentTransformationConfig with partial parameters merged with defaults.
524 """
525 params = EquivalentTransformationConfig.DEFAULTS.copy()
526 params.update(kwargs)
528 **params, optional=EquivalentTransformationConfig.optional
529 )
530
531
532# }@
from_dict(cls, dict config_dict)
Create EquivalentTransformationConfig from dictionary (JSON structure)
check_valid(self, except_list=["optional"])
Override to skip py_to_cpp_name check (uses nested structure in to_dict)
EquivalentTransformationConfig get_equivalent_transformation_config(**kwargs)
Create EquivalentTransformationConfig with partial parameters merged with defaults.
__init__(self, int seed, Dict[str, str] layer_name_pattern, bool apply_hadamard_rotation_matrix, bool norm_conv_apply, bool norm_conv_learn, float norm_conv_smoothing_factor, float norm_conv_min_gamma, float norm_conv_max_gamma, bool qk_apply, float qk_smoothing_factor, float qk_min_gamma, float qk_max_gamma, bool ud_apply, bool ud_learn, float ud_smoothing_factor, float ud_min_gamma, float ud_max_gamma, bool vo_apply, float vo_smoothing_factor, float vo_min_gamma, float vo_max_gamma, bool lut_rescale_apply, bool lut_rescale_alpha, bool lut_rescale_beta, bool ff_multi_lut_apply, List[float] ff_multi_lut_breakpoints, bool spin_r1_apply, str spin_r1_matrix_path, bool head_out_ch_rotation_apply, str head_out_ch_rotation_matrix_path, bool spin_r2_apply, bool spin_r2_learn, str spin_r2_matrix_path, bool qk_rotation_apply, str qk_rotation_matrix_path, bool flatten_quant_apply, bool flatten_quant_learn, float flatten_quant_apply_threshold, float flatten_quant_max_overhead, bool split_ffn_apply, int split_ffn_ch_per_ffn, bool optional=optional)
Initialize the EquivalentTransformationConfig.