Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
tensor_quant_overrides.py521 linesDownload Raw Back to quantization
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 
codekingpro/portable-devtools · Team Ai