Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
qdq_loss_debug.py390 linesDownload Raw Back to quantization
1# --------------------------------------------------------------------------
2# Copyright (c) Microsoft, Intel Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6
7"""Utilities to run a given ONNX model, while saving input/output tensors of
8eligible operator nodes.
9
10A use case is to debug quantization induced accuracy drop. An AI engineer can
11run the original float32 model and the quantized model with the same inputs,
12then compare the corresponding activations between the two models to find
13where the divergence is.
14
15Example Usage:
16
17```python
18    class ExampleDataReader(CalibrationDataReader):
19        def __init__(self):
20            ...
21        def get_next(self):
22            ...
23
24    input_data_reader = ExampleDataReader()
25
26    augmented_model_path = str(Path(self._tmp_model_dir.name).joinpath("augmented_model.onnx"))
27    modify_model_output_intermediate_tensors (path_to_onnx_model, augmented_model_path)
28
29    tensor_dict = collect_activations(augmented_model_path, input_data_reader)
30```
31
32`tensor_dict` points to a dictionary where the keys are tensor names and each value
33is a list of tensors, one from each model run
34
35"""
36
37import logging
38import math
39import time
40from collections.abc import Callable, Sequence
41from pathlib import Path
42
43import numpy
44import onnx
45from onnx import helper, numpy_helper
46
47import onnxruntime
48
49from .calibrate import CalibraterBase, CalibrationDataReader
50from .onnx_model import ONNXModel
51from .quant_utils import (
52    DEQUANT_OP_NAME,
53    DEQUANT_OUTPUT_SUFFIX,
54    QUANT_INPUT_SUFFIX,
55    TENSOR_NAME_QUANT_SUFFIX,
56    find_by_name,
57    load_model_with_shape_infer,
58)
59
60_TENSOR_SAVE_POSTFIX = "_ReshapedSavedOutput"
61_TENSOR_SAVE_POSTFIX_LEN = len(_TENSOR_SAVE_POSTFIX)
62
63
64def modify_model_output_intermediate_tensors(
65    input_model_path: str | Path,
66    output_model_path: str | Path,
67    op_types_for_saving: Sequence[str] | None = None,
68    save_as_external_data: bool = False,
69) -> None:
70    """Augment a given ONNX model to save node input/output tensors.
71
72    Add all input/output tensors of operator nodes to model outputs
73    so that their values can be retrieved for debugging purposes.
74
75    Args:
76        input_model: the path to load the model.
77        op_types_for_saving: Operator types for which the
78                input/output should be saved. By default, saving all the
79                float32/float16 tensors.
80
81    Returns:
82        The augmented ONNX model
83    """
84
85    if op_types_for_saving is None:
86        op_types_for_saving = []
87    saver = CalibraterBase(input_model_path, op_types_to_calibrate=op_types_for_saving)
88    model_to_augment = saver.model
89    tensors, value_infos = saver.select_tensors_to_calibrate(model_to_augment)
90    reshape_shape_name = "LinearReshape_" + str(time.time())
91    reshape_shape = numpy_helper.from_array(numpy.array([-1], dtype=numpy.int64), reshape_shape_name)
92    model_to_augment.graph.initializer.append(reshape_shape)
93
94    for tensor_name in tensors:
95        reshape_output = tensor_name + _TENSOR_SAVE_POSTFIX
96        reshape_node = onnx.helper.make_node(
97            "Reshape",
98            inputs=[tensor_name, reshape_shape_name],
99            outputs=[reshape_output],
100            name=reshape_output,
101        )
102        model_to_augment.graph.node.append(reshape_node)
103        reshape_output_value_info = helper.make_tensor_value_info(
104            reshape_output, value_infos[tensor_name].type.tensor_type.elem_type, [-1]
105        )
106        model_to_augment.graph.output.append(reshape_output_value_info)
107
108    onnx.save(
109        model_to_augment,
110        output_model_path,
111        save_as_external_data=save_as_external_data,
112    )
113
114
115def collect_activations(
116    augmented_model: str,
117    input_reader: CalibrationDataReader,
118    session_options=None,
119    execution_providers: Sequence[str] | None = None,
120) -> dict[str, list[numpy.ndarray]]:
121    """Run augmented model and collect activations tensors.
122
123    Args:
124        augmented_model: Path to augmented model created by modify_model_output_intermediate_tensors ()
125        input_reader: Logic for reading input for the model, augmented model have the same
126            input with the original model.
127        session_options: Optional OnnxRuntime session options for controlling model run.
128            By default graph optimization is turned off
129        execution_providers: Collection of execution providers for running the model.
130            Only CPU EP is used by default.
131
132    Returns:
133        A dictionary where the key is tensor name and values are list of tensors from each batch
134    """
135
136    if session_options is None:
137        session_options = onnxruntime.SessionOptions()
138        session_options.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL
139    if execution_providers is None:
140        execution_providers = ["CPUExecutionProvider"]
141
142    inference_session = onnxruntime.InferenceSession(
143        augmented_model,
144        sess_options=session_options,
145        providers=execution_providers,
146    )
147
148    intermediate_outputs = []
149    for input_d in input_reader:
150        intermediate_outputs.append(inference_session.run(None, input_d))
151    if not intermediate_outputs:
152        raise RuntimeError("No data is collected while running augmented model!")
153
154    output_dict = {}
155    output_info = inference_session.get_outputs()
156    for batch in intermediate_outputs:
157        for output, output_data in zip(output_info, batch, strict=False):
158            if output.name.endswith(_TENSOR_SAVE_POSTFIX):
159                output_name = output.name[:-_TENSOR_SAVE_POSTFIX_LEN]
160                output_dict.setdefault(output_name, []).append(output_data)
161
162    return output_dict
163
164
165_POST_QDQ_POSTFIX1 = DEQUANT_OUTPUT_SUFFIX + "_1"
166
167
168def _add_pre_post_qdq_pair(
169    qdq_cmp: dict[str, dict[str, Sequence[numpy.ndarray]]],
170    activation_name: str,
171    pre_qdq_tensors: Sequence[numpy.ndarray] | None,
172    post_qdq_tensors: Sequence[numpy.ndarray] | None,
173) -> None:
174    if post_qdq_tensors is not None and pre_qdq_tensors is not None:
175        qdq_cmp[activation_name] = {}
176        qdq_cmp[activation_name]["pre_qdq"] = pre_qdq_tensors
177        qdq_cmp[activation_name]["post_qdq"] = post_qdq_tensors
178
179
180def create_activation_matching(
181    qdq_activations: dict[str, Sequence[numpy.ndarray]],
182    float_activations: dict[str, Sequence[numpy.ndarray]] | None = None,
183) -> dict[str, dict[str, Sequence[numpy.ndarray]]]:
184    """Comparing activation values to help debugging accuracy loss due to quantization.
185
186    This functions takes saved activations from the QDQ model and (optionally) the
187    float point model, and provides a data structure for comparing:
188        * from the qdq model, activation values before and after QDQ operation
189        * across both models, activations from the orignal model vs the corresponding
190          activations in the QDQ model
191
192    Arg:
193        qdq_activations: Output of `collect_activations`. This must be from a quantized
194            model with QDQ format.
195        float_activations: Output of `collect_activations`. This must be from the float
196            point model.
197
198    Returns:
199        Dict for comparing pre and post quantized activation tensors. E.g.
200        ```
201        qdq_cmp = cmp_qdq_input_output(qdq_activations)
202        print(qdq_cmp['activation1']['pre_qdq'][0])
203        print(qdq_cmp['activation1'][`post_qdq'][0])
204
205
206        qdq_cmp = cmp_qdq_input_output(qdq_activations, float_activations)
207        print(qdq_cmp['activation1']['float'][0])
208        print(qdq_cmp['activation1']['pre_qdq'][0])
209        print(qdq_cmp['activation1'][`post_qdq'][0])
210        ```
211    """
212
213    qdq_cmp: dict[str, dict[str, Sequence[numpy.ndarray]]] = {}
214    for tensor_name, tensors in qdq_activations.items():
215        if tensor_name.endswith(QUANT_INPUT_SUFFIX):
216            pre_name = tensor_name[: -len(QUANT_INPUT_SUFFIX)]
217            post_qdq_tensors = qdq_activations.get(pre_name)
218            pre_qdq_tensors = tensors
219            _add_pre_post_qdq_pair(qdq_cmp, pre_name, pre_qdq_tensors, post_qdq_tensors)
220        elif tensor_name.endswith(DEQUANT_OUTPUT_SUFFIX):
221            pre_name = tensor_name[: -len(DEQUANT_OUTPUT_SUFFIX)]
222            pre_qdq_tensors = qdq_activations.get(pre_name)
223            post_qdq_tensors = tensors
224            _add_pre_post_qdq_pair(qdq_cmp, pre_name, pre_qdq_tensors, post_qdq_tensors)
225        elif tensor_name.endswith(_POST_QDQ_POSTFIX1):
226            pre_name = tensor_name[: -len(_POST_QDQ_POSTFIX1)]
227            pre_qdq_tensors = qdq_activations.get(pre_name)
228            post_qdq_tensors = tensors
229            _add_pre_post_qdq_pair(qdq_cmp, pre_name, pre_qdq_tensors, post_qdq_tensors)
230
231    if not float_activations:
232        return qdq_cmp
233
234    for act_name, act_values in qdq_cmp.items():
235        float_acts = float_activations.get(act_name)
236        if float_acts is not None:
237            act_values["float"] = float_acts
238
239    return qdq_cmp
240
241
242def _run_dequantize_linear(
243    weight_tensor: numpy.ndarray, weight_scale: numpy.ndarray, weight_zp: numpy.ndarray, channel_axis: int
244) -> numpy.ndarray | None:
245    assert weight_scale.shape == weight_zp.shape
246    if weight_zp.size == 1:
247        return (weight_tensor - weight_zp) * weight_scale
248
249    assert weight_zp.ndim == 1
250    reshape_dims = list(weight_tensor.shape)  # deep copy
251    reshape_dims[channel_axis] = 1  # only one per channel for reshape
252    channel_count = weight_tensor.shape[channel_axis]
253    dequantized_weights = None
254    for i in range(channel_count):
255        per_channel_data = weight_tensor.take(i, channel_axis)
256        dequantized_per_channel_data = (per_channel_data - weight_zp[i]) * weight_scale[i]
257        if i == 0:
258            dequantized_weights = numpy.asarray(dequantized_per_channel_data).reshape(reshape_dims)
259        else:
260            channel_weights = numpy.asarray(dequantized_per_channel_data).reshape(reshape_dims)
261            dequantized_weights = numpy.concatenate((dequantized_weights, channel_weights), channel_axis)
262
263    if dequantized_weights is None:
264        return None
265
266    dequantized_weights.reshape(weight_tensor.shape)
267    return dequantized_weights
268
269
270def create_weight_matching(float_model_path: str, qdq_model_path: str) -> dict[str, dict[str, numpy.ndarray]]:
271    """Comparing weight values to help debugging accuracy loss due to quantization.
272
273    This functions takes the float model and the qdq model, and provides a data structure for comparing
274    their corresponding weights to locate quantization errors
275
276    Arg:
277        float_model_path: Path points to the float point model.
278        qdq_model_path: Path points to the qdq model.
279
280    Returns:
281        Dict for comparing weight tensors. E.g.
282        ```
283        qdq_weight_cmp = create_weight_matching(float_model, qdq_model)
284        print(qdq_weight_cmp['activation1']['float'])
285        print(qdq_weight_cmp['activation1']['dequantized'])
286        ```
287    """
288    float_onnx_model = ONNXModel(load_model_with_shape_infer(Path(float_model_path)))
289    qdq_onnx_model = ONNXModel(load_model_with_shape_infer(Path(qdq_model_path)))
290
291    matched_weights: dict[str, dict[str, numpy.ndarray]] = {}
292    initializers = qdq_onnx_model.initializer()
293    for node in qdq_onnx_model.nodes():
294        if node.op_type != DEQUANT_OP_NAME:
295            continue  # Only care about DQ node
296        weight_name: str = node.input[0]
297        weight_values = find_by_name(weight_name, initializers)
298        if not weight_values:
299            continue  # Only care about DQ node with const inputs
300        if not weight_name.endswith(TENSOR_NAME_QUANT_SUFFIX):
301            logging.error(f"Model Error in '{qdq_model_path}': Dequantized tensor name '{weight_name}' not recognized!")
302            continue
303
304        axis = -1
305        for attr in node.attribute:
306            if attr.name == "axis":
307                axis = attr.i
308
309        weight_tensor = numpy_helper.to_array(weight_values)
310        weight_scale = numpy_helper.to_array(find_by_name(node.input[1], initializers))
311        if len(node.input) > 2:
312            weight_zp = numpy_helper.to_array(find_by_name(node.input[2], initializers))
313        else:
314            weight_zp = numpy.zeros(weight_scale.shape, dtype=numpy.int32)
315
316        # Perform dequantization:
317        if weight_scale.size == weight_zp.size == 1:
318            # Avoids the confusion between a scaler and a tensor of one element.
319            weight_scale = weight_scale.reshape(())
320            weight_zp = weight_zp.reshape(())
321        if weight_scale.shape != weight_zp.shape:
322            raise RuntimeError(
323                f"scale and zero_point must have the same shape but {weight_scale.shape} != {weight_zp.shape}"
324            )
325        weight_quant = _run_dequantize_linear(weight_tensor, weight_scale, weight_zp, channel_axis=axis)
326        weight_name = weight_name[: -len(TENSOR_NAME_QUANT_SUFFIX)]
327        if weight_quant is None:
328            logging.error(f"Model Error in '{qdq_model_path}': '{weight_name}' per-channel quantization on 0 channel")
329            continue
330
331        float_values = find_by_name(weight_name, float_onnx_model.initializer())
332        if not float_values:
333            logging.error(f"Model Error in '{float_model_path}': weight tensor '{weight_name}' not found!")
334            continue
335        weight_float = numpy_helper.to_array(float_values)
336        matched_weights[weight_name] = {"float": weight_float, "dequantized": weight_quant}
337
338    return matched_weights
339
340
341def compute_signal_to_quantization_noice_ratio(
342    x: Sequence[numpy.ndarray] | numpy.ndarray, y: Sequence[numpy.ndarray] | numpy.ndarray
343) -> float:
344    if isinstance(x, numpy.ndarray):
345        xlist = [x]
346    else:
347        xlist = x
348    if isinstance(y, numpy.ndarray):
349        ylist = [y]
350    else:
351        ylist = y
352    if len(xlist) != len(ylist):
353        raise RuntimeError("Unequal number of tensors to compare!")
354
355    left = numpy.concatenate(xlist).flatten()
356    right = numpy.concatenate(ylist).flatten()
357
358    epsilon = numpy.finfo("float").eps
359    tensor_norm = max(numpy.linalg.norm(left), epsilon)
360    diff_norm = max(numpy.linalg.norm(left - right), epsilon)
361    res = tensor_norm / diff_norm
362    return 20 * math.log10(res)
363
364
365def compute_weight_error(
366    weights_match: dict[str, dict[str, numpy.ndarray]],
367    err_func: Callable[[numpy.ndarray, numpy.ndarray], float] = compute_signal_to_quantization_noice_ratio,
368) -> dict[str, float]:
369    result: dict[str, float] = {}
370    for weight_name, weight_match in weights_match.items():
371        result[weight_name] = err_func(weight_match["float"], weight_match["dequantized"])
372    return result
373
374
375def compute_activation_error(
376    activations_match: dict[str, dict[str, Sequence[numpy.ndarray]]],
377    err_func: Callable[
378        [Sequence[numpy.ndarray], Sequence[numpy.ndarray]], float
379    ] = compute_signal_to_quantization_noice_ratio,
380) -> dict[str, dict[str, float]]:
381    result: dict[str, dict[str, float]] = {}
382    for name, match in activations_match.items():
383        err_result: dict[str, float] = {}
384        err_result["qdq_err"] = err_func(match["pre_qdq"], match["post_qdq"])
385        float_activation = match["float"]
386        if float_activation:
387            err_result["xmodel_err"] = err_func(float_activation, match["post_qdq"])
388        result[name] = err_result
389    return result
390 
codekingpro/portable-devtools · Team Ai