codekingpro/portable-devtools
115k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6from __future__ import annotations
7
8import json
9from collections.abc import MutableMapping
10from dataclasses import dataclass
11from typing import Any
12
13import onnx
14
15from .quant_utils import QuantType
16
17
18@dataclass
19class QuantTypeInfo: # noqa: PLW1641
20 """
21 The quantization type information for a tensor override.
22 """
23
24 quant_type: QuantType
25 symmetric: bool | None = None # If None, assumes default is used.
26 reduce_range: bool | None = None # If None, assumes default is used.
27 axis: int | None = None # If None, assumes per-tensor quantization
28
29 def __eq__(self, other: object):
30 if isinstance(other, QuantTypeInfo):
31 return (
32 self.quant_type == other.quant_type
33 and (self.symmetric is None or other.symmetric is None or self.symmetric == other.symmetric)
34 and (self.reduce_range is None or other.reduce_range is None or self.reduce_range == other.reduce_range)
35 and (self.axis == other.axis)
36 )
37 return NotImplemented
38
39 @staticmethod
40 def load_from_dict(
41 raw_dict: dict[str, Any],
42 default_qtype: QuantType | None = None,
43 default_symmetric: bool | None = None,
44 default_reduce_range: bool | None = None,
45 ) -> QuantTypeInfo:
46 return QuantTypeInfo(
47 raw_dict.get("quant_type", default_qtype),
48 raw_dict.get("symmetric", default_symmetric),
49 raw_dict.get("reduce_range", default_reduce_range),
50 raw_dict.get("axis"),
51 )
52
53 def save_to_dict(self, raw_dict: dict[str, Any]):
54 raw_dict["quant_type"] = self.quant_type
55 if self.symmetric is not None:
56 raw_dict["symmetric"] = self.symmetric
57 if self.reduce_range is not None:
58 raw_dict["reduce_range"] = self.reduce_range
59 if self.axis is not None:
60 raw_dict["axis"] = self.axis
61
62
63class TensorQuantOverridesHelper(MutableMapping):
64 """
65 Utility wrapper over the tensor quantization overrides passed via extra_options.
66 """
67
68 def __init__(self, raw_overrides: dict[str, list[dict[str, Any]]]):
69 self.overrides = raw_overrides
70 self.quant_types = None
71 self.keys_unsupported_with_scale_zp = {"symmetric", "reduce_range", "rmax", "rmin"}
72
73 def has_per_tensor_overrides(self, tensor_name: str) -> bool:
74 overrides_list = self.overrides.get(tensor_name)
75 return overrides_list and "axis" not in overrides_list[0]
76
77 def has_per_channel_overrides(self, tensor_name: str) -> bool:
78 overrides_list = self.overrides.get(tensor_name)
79 return overrides_list and "axis" in overrides_list[0]
80
81 def overrides_scale_zp(self, tensor_name: str) -> bool:
82 overrides_list = self.overrides.get(tensor_name)
83 return overrides_list and ("scale" in overrides_list[0]) and ("zero_point" in overrides_list[0])
84
85 def get_per_tensor_overrides(
86 self,
87 tensor_name: str,
88 default_val: dict[str, Any] | None = None,
89 ) -> dict[str, Any] | None:
90 default_list_val = [default_val] if default_val is not None else None
91 overrides_list = self.overrides.get(tensor_name, default_list_val)
92 if overrides_list and "axis" in overrides_list[0]:
93 raise ValueError(
94 f"Expected tensor '{tensor_name}' to use per-tensor quantization overrides, "
95 f"but found per-channel overrides."
96 )
97
98 return overrides_list[0] if overrides_list else None
99
100 def get_per_channel_overrides(
101 self,
102 tensor_name: str,
103 default_val: list[dict[str, Any]] | None = None,
104 ) -> list[dict[str, Any]] | None:
105 overrides_list = self.overrides.get(tensor_name, default_val)
106
107 if not overrides_list:
108 return None
109
110 if "axis" not in overrides_list[0]:
111 raise ValueError(
112 f"Expected tensor '{tensor_name}' to have per-channel quantization overrides (axis value is missing).",
113 )
114
115 return overrides_list
116
117 def get_quant_types(self) -> set[QuantType]:
118 if self.quant_types is not None:
119 return self.quant_types
120
121 self.quant_types = set()
122
123 if self.overrides:
124 for quant_overrides_list in self.overrides.values():
125 for quant_overrides in quant_overrides_list:
126 if "quant_type" in quant_overrides:
127 self.quant_types.add(quant_overrides["quant_type"])
128
129 if "convert" in quant_overrides and "quant_type" in quant_overrides["convert"]:
130 self.quant_types.add(quant_overrides["convert"]["quant_type"])
131
132 return self.quant_types
133
134 def _is_valid_per_tensor(
135 self,
136 initializers,
137 default_activation_qtype,
138 tensor_name: str,
139 quant_overrides: dict[str, Any],
140 ) -> tuple[bool, str | None]:
141 if not isinstance(quant_overrides, dict):
142 return (
143 False,
144 f"Tensor quantization overrides for '{tensor_name}' are not in a dict",
145 )
146
147 is_initializer = tensor_name in initializers
148
149 quant_type = quant_overrides.get("quant_type")
150 if quant_type:
151 self.quant_types.add(quant_type)
152
153 has_scale = "scale" in quant_overrides
154 has_zero_point = "zero_point" in quant_overrides
155
156 if (has_scale and not has_zero_point) or (has_zero_point and not has_scale):
157 return (
158 False,
159 "Must provide both 'scale' and 'zero_point' if one of the overrides is provided",
160 )
161
162 if has_scale:
163 keys = self.keys_unsupported_with_scale_zp.intersection(set(quant_overrides))
164 if keys:
165 return (
166 False,
167 f"Tensor override option(s) [{', '.join(keys)}] are invalid with 'scale' and 'zero_point'",
168 )
169
170 if "reduce_range" in quant_overrides and not is_initializer:
171 return (
172 False,
173 f"Option 'reduce_range' is only supported for initializers, not for activation {tensor_name}",
174 )
175
176 if "convert" in quant_overrides:
177 if is_initializer:
178 return False, "Cannot use 'convert' override for initializers"
179
180 if "quant_type" not in quant_overrides["convert"]:
181 return False, f"'convert' options (tensor '{tensor_name}') must specify a 'quant_type'"
182
183 if "reduce_range" in quant_overrides["convert"]:
184 return (
185 False,
186 f"Option 'reduce_range' is only supported for initializers, not for activation {tensor_name}",
187 )
188
189 convert_quant_type = quant_overrides["convert"]["quant_type"]
190 original_quant_type = quant_type if quant_type is not None else default_activation_qtype
191 if convert_quant_type == original_quant_type:
192 return (
193 False,
194 f"'convert' quant_type must differ from original quant_type (tensor '{tensor_name}')",
195 )
196
197 convert_has_scale = "scale" in quant_overrides["convert"]
198 convert_has_zero_point = "zero_point" in quant_overrides["convert"]
199
200 if (convert_has_scale and not convert_has_zero_point) or (convert_has_zero_point and not convert_has_scale):
201 return (
202 False,
203 f"Must provide both 'scale' and 'zero_point' if one of the overrides is provided (tensor '{tensor_name}')",
204 )
205
206 if convert_has_scale:
207 keys = self.keys_unsupported_with_scale_zp.intersection(set(quant_overrides["convert"]))
208 if keys:
209 return (
210 False,
211 f"Tensor override option(s) [{', '.join(keys)}] are invalid with 'scale' and 'zero_point' "
212 f"(tensor '{tensor_name}')",
213 )
214
215 self.quant_types.add(convert_quant_type)
216
217 return True, None
218
219 def _is_valid_per_channel(
220 self,
221 initializers,
222 tensor_name: str,
223 quant_overrides_list: list[dict[str, Any]],
224 ) -> tuple[bool, str | None]:
225 is_initializer = tensor_name in initializers
226
227 if not is_initializer:
228 return (
229 False,
230 f"Tensor '{tensor_name}' has per-channel overrides, but is not an initializer",
231 )
232
233 axis = quant_overrides_list[0].get("axis")
234
235 if axis is None:
236 return (
237 False,
238 f"Per-channel overrides for tensor {tensor_name} is missing an 'axis' value in "
239 "the first channel dictionary.",
240 )
241
242 weight_shape = list(initializers[tensor_name].dims)
243 weight_rank = len(weight_shape)
244 norm_axis = axis
245 if norm_axis < 0:
246 norm_axis += weight_rank
247
248 if norm_axis < 0 or norm_axis >= len(weight_shape):
249 return (
250 False,
251 f"Axis override value is out-of-bounds for tensor {tensor_name} (rank {len(weight_shape)})",
252 )
253
254 if len(quant_overrides_list) > 1 and len(quant_overrides_list) != weight_shape[norm_axis]:
255 return (
256 False,
257 f"Incorrect number of channel overrides for tensor {tensor_name} (axis {axis}), "
258 f"expected {weight_shape[axis]}, but found {len(quant_overrides_list)}.",
259 )
260
261 if "convert" in quant_overrides_list[0]:
262 return False, f"Cannot use 'convert' override for initializers, such as {tensor_name}."
263
264 quant_type = quant_overrides_list[0].get("quant_type")
265 if quant_type:
266 self.quant_types.add(quant_type)
267
268 symmetric = quant_overrides_list[0].get("symmetric")
269 reduce_range = quant_overrides_list[0].get("reduce_range")
270
271 has_scale = "scale" in quant_overrides_list[0]
272 has_zero_point = "zero_point" in quant_overrides_list[0]
273 has_scale_zp = has_scale and has_zero_point
274
275 if (has_scale and not has_zero_point) or (has_zero_point and not has_scale):
276 return (
277 False,
278 "Must provide both 'scale' and 'zero_point' if one of the overrides is provided",
279 )
280
281 if has_scale_zp:
282 keys = self.keys_unsupported_with_scale_zp.intersection(set(quant_overrides_list[0]))
283 if keys:
284 return (
285 False,
286 f"Tensor override option(s) [{', '.join(keys)}] are invalid with 'scale' and 'zero_point'",
287 )
288
289 has_rmin = "rmin" in quant_overrides_list[0]
290 has_rmax = "rmax" in quant_overrides_list[0]
291 has_rmin_rmax = has_rmin and has_rmax
292 if (has_rmin and not has_rmax) or (not has_rmin and has_rmax):
293 return (
294 False,
295 "Must provide both 'rmin' and 'rmax' if one is provided",
296 )
297
298 for index, quant_overrides in enumerate(quant_overrides_list[1:]):
299 if not isinstance(quant_overrides, dict):
300 return (
301 False,
302 f"Tensor quantization overrides at index {index} for '{tensor_name}' are not in a dict",
303 )
304
305 if "convert" in quant_overrides:
306 return False, f"Cannot use 'convert' override for initializers, such as {tensor_name}."
307
308 # For per-channel quantization, all channels must use the same quantization type, axis, symmetric
309 # and reduce_range values. And, if specified, they must be present in the first channel dict
310 # (i.e., quant_overrides_list[0]).
311 if "quant_type" in quant_overrides and quant_type != quant_overrides["quant_type"]:
312 return (
313 False,
314 "Channel quantization types for tensor '{tensor_name}' do not match at index {index}.",
315 )
316 if "axis" in quant_overrides and axis != quant_overrides["axis"] and norm_axis != quant_overrides["axis"]:
317 return (
318 False,
319 "Channel axis for tensor '{tensor_name}' does not match at index {index}.",
320 )
321 if "symmetric" in quant_overrides and symmetric != quant_overrides["symmetric"]:
322 return (
323 False,
324 "Channel symmetric value for tensor '{tensor_name}' does not match at index {index}.",
325 )
326 if "reduce_range" in quant_overrides and reduce_range != quant_overrides["reduce_range"]:
327 return (
328 False,
329 "Channel reduce_range value for tensor '{tensor_name}' does not match at index {index}.",
330 )
331
332 # If override scale/zp, must do so for all channels.
333 chan_has_scale_zp = "scale" in quant_overrides and "zero_point" in quant_overrides
334
335 if has_scale_zp and not chan_has_scale_zp:
336 return (
337 False,
338 "Per-channel overrides that specify scale/zero_point must do so for all channels, "
339 f"but tensor '{tensor_name}' is missing them at index {index}.",
340 )
341
342 if chan_has_scale_zp:
343 keys = self.keys_unsupported_with_scale_zp.intersection(set(quant_overrides))
344 if keys:
345 return (
346 False,
347 f"Tensor override option(s) [{', '.join(keys)}] are invalid with 'scale' and 'zero_point'",
348 )
349
350 # If override rmin/rmax, must do so for all channels.
351 chan_has_rmin_rmax = "rmin" in quant_overrides and "rmax" in quant_overrides
352 if has_rmin_rmax and not chan_has_rmin_rmax:
353 return (
354 False,
355 "Per-channel overrides that specify rmin/rmax must do so for all channels, "
356 f"but tensor '{tensor_name}' is missing them at index {index}.",
357 )
358
359 return True, None
360
361 def is_valid(
362 self,
363 initializers: dict[str, onnx.TensorProto],
364 activation_names: set[str],
365 default_activation_qtype,
366 ) -> tuple[bool, str | None]:
367 self.quant_types = set()
368
369 # Validate that compatible/valid overrides are provided.
370 if self.overrides:
371 for tensor_name, quant_overrides_list in self.overrides.items():
372 if tensor_name not in initializers and tensor_name not in activation_names:
373 return False, f"Tensor '{tensor_name}' in TensorQuantOverrides is not present in the model"
374
375 if not isinstance(quant_overrides_list, list):
376 return False, f"Tensor quantization overrides for '{tensor_name}' are not in a list"
377
378 if not quant_overrides_list:
379 continue
380
381 if not isinstance(quant_overrides_list[0], dict):
382 return False, f"Tensor quantization overrides at index 0 for '{tensor_name}' are not in a dict"
383
384 if not quant_overrides_list[0]:
385 continue
386
387 axis = quant_overrides_list[0].get("axis")
388 is_per_channel = len(quant_overrides_list) > 1 or axis is not None
389
390 if is_per_channel:
391 return self._is_valid_per_channel(initializers, tensor_name, quant_overrides_list)
392
393 return self._is_valid_per_tensor(
394 initializers, default_activation_qtype, tensor_name, quant_overrides_list[0]
395 )
396
397 return True, None
398
399 def update_tensor_overrides(
400 self,
401 tensor_name: str,
402 new_vals: dict[str, Any],
403 channels: list[int] | None = None,
404 overwrite: bool = True,
405 ) -> bool:
406 if not new_vals:
407 return False
408
409 channels = set(channels) if channels is not None else None
410 have_overrides = self.overrides.get(tensor_name)
411
412 # If `overwrite` is False, check if we would overwrite anything.
413 do_update = True
414 if not overwrite and have_overrides:
415 for channel, overrides in enumerate(self.overrides[tensor_name]):
416 if channels is not None and channel not in channels:
417 continue
418 if set(new_vals).intersection(set(overrides)):
419 do_update = False
420 break
421
422 # Do the update if `overwrite` is True or if nothing is overwritten (do not want partial overwrites).
423 if do_update:
424 if not have_overrides:
425 self.overrides[tensor_name] = [{}]
426
427 for channel, overrides in enumerate(self.overrides[tensor_name]):
428 if channels is not None and channel not in channels:
429 continue
430 overrides.update(new_vals)
431
432 return do_update
433
434 def get_node_output_qtype_info(
435 self,
436 output_name: str,
437 default_qtype: QuantType | None,
438 default_symmetric: bool | None = None,
439 ) -> QuantTypeInfo:
440 # Outputs are activations, which do not support 'reduce_range' or 'axis'
441 if output_name not in self.overrides:
442 return QuantTypeInfo(default_qtype, default_symmetric)
443
444 tensor_overrides = self.overrides[output_name][0]
445
446 return QuantTypeInfo(
447 tensor_overrides.get("quant_type", default_qtype),
448 tensor_overrides.get("symmetric", default_symmetric),
449 )
450
451 def get_node_input_qtype_info(
452 self,
453 input_name: str,
454 node_name: str,
455 default_qtype: QuantType | None,
456 default_symmetric: bool | None = None,
457 default_reduce_range: bool | None = None,
458 ) -> QuantTypeInfo:
459 if input_name not in self.overrides or not self.overrides[input_name]:
460 return QuantTypeInfo(default_qtype, default_symmetric, default_reduce_range)
461
462 # Get the first overrides dict in the list. This works for both per-tensor and per-channel
463 # quantization because all channels must use the same quant type.
464 tensor_overrides = self.overrides[input_name][0]
465 producer_type = tensor_overrides.get("quant_type", default_qtype)
466
467 if "convert" not in tensor_overrides:
468 return QuantTypeInfo(
469 producer_type,
470 tensor_overrides.get("symmetric", default_symmetric),
471 tensor_overrides.get("reduce_range", default_reduce_range),
472 tensor_overrides.get("axis"),
473 )
474
475 # This tensor is converted. Check if the node gets the original qtype or the converted qtype.
476 convert_dict = tensor_overrides["convert"]
477 qtype_info = QuantTypeInfo(
478 producer_type,
479 convert_dict.get("symmetric", default_symmetric),
480 # Converted tensors are not initializers, so do not have 'axis' or 'reduce_range'.
481 )
482
483 # Check if all nodes receive the converted type (i.e., recv_nodes is None) or this node
484 # is in the list of consumers (recv_nodes).
485 if ("recv_nodes" not in convert_dict) or (node_name in convert_dict["recv_nodes"]):
486 qtype_info.quant_type = convert_dict["quant_type"]
487
488 return qtype_info
489
490 def pprint_str(self, indent=None) -> str:
491 return json.dumps(self.overrides, default=str, indent=indent)
492
493 def empty(self) -> bool:
494 return not self.overrides
495
496 def get_dict(self) -> dict[str, list[dict[str, Any]]]:
497 return self.overrides
498
499 # Required implementations of abstract methods in collections.abc.MutableMapping
500 # so that this class can be used like a dict.
501 def __setitem__(self, key: str, value: list[dict]):
502 self.overrides[key] = value
503
504 def __getitem__(self, key: str) -> list[dict]:
505 return self.overrides[key]
506
507 def __delitem__(self, key: str):
508 del self.overrides[key]
509
510 def __iter__(self):
511 return iter(self.overrides)
512
513 def __len__(self):
514 return len(self.overrides)
515
516 def __str__(self) -> str:
517 return str(self.overrides)
518
519 def __repr__(self) -> str:
520 return f"{super().__repr__()}, TensorQuantOverridesHelper({self.overrides})"
521 