codekingpro/portable-devtools
114k
1#
2# The implementation of this file is based on:
3# https://github.com/intel/neural-compressor/tree/master/neural_compressor
4#
5# Copyright (c) 2023 Intel Corporation
6#
7# Licensed under the Apache License, Version 2.0 (the "License");
8# you may not use this file except in compliance with the License.
9# You may obtain a copy of the License at
10#
11# http://www.apache.org/licenses/LICENSE-2.0
12#
13# Unless required by applicable law or agreed to in writing, software
14# distributed under the License is distributed on an "AS IS" BASIS,
15# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
16# See the License for the specific language governing permissions and
17# limitations under the License.
18
19"""Helper classes or functions for onnxrt adaptor."""
20
21import importlib
22import logging
23
24import numpy as np
25
26logger = logging.getLogger("neural_compressor")
27
28
29MAXIMUM_PROTOBUF = 2147483648
30
31
32def simple_progress_bar(total, i):
33 """Progress bar for cases where tqdm can't be used."""
34 progress = i / total
35 bar_length = 20
36 bar = "#" * int(bar_length * progress)
37 spaces = " " * (bar_length - len(bar))
38 percentage = progress * 100
39 print(f"\rProgress: [{bar}{spaces}] {percentage:.2f}%", end="")
40
41
42def find_by_name(name, item_list):
43 """Helper function to find item by name in a list."""
44 items = []
45 for item in item_list:
46 assert hasattr(item, "name"), f"{item} should have a 'name' attribute defined" # pragma: no cover
47 if item.name == name:
48 items.append(item)
49 if len(items) > 0:
50 return items[0]
51 else:
52 return None
53
54
55def to_numpy(data):
56 """Convert to numpy ndarrays."""
57 import torch # noqa: PLC0415
58
59 if not isinstance(data, np.ndarray):
60 if not importlib.util.find_spec("torch"):
61 logger.error(
62 "Please install torch to enable subsequent data type check and conversion, "
63 "or reorganize your data format to numpy array."
64 )
65 exit(0)
66 if isinstance(data, torch.Tensor):
67 if data.dtype is torch.bfloat16: # pragma: no cover
68 return data.detach().cpu().to(torch.float32).numpy()
69 if data.dtype is torch.chalf: # pragma: no cover
70 return data.detach().cpu().to(torch.cfloat).numpy()
71 return data.detach().cpu().numpy()
72 else:
73 try:
74 return np.array(data)
75 except Exception:
76 assert False, ( # noqa: B011
77 f"The input data for onnx model is {type(data)}, which is not supported to convert to numpy ndarrays."
78 )
79 else:
80 return data
81 