codekingpro/portable-devtools
115k
1class QuantOperatorBase:
2 def __init__(self, onnx_quantizer, onnx_node):
3 self.quantizer = onnx_quantizer
4 self.node = onnx_node
5
6 def should_quantize(self):
7 if not self.quantizer.should_quantize_node(self.node):
8 return False
9
10 return self.quantizer.is_float_tensor(self.node.input[0])
11
12 def quantize(self):
13 """
14 Given a node which does not support quantization, this method checks whether the input to
15 this node is quantized and adds a DequantizeLinear node to dequantize this input back to FP32
16 parameter node: Current node
17 parameter new_nodes_list: List of new nodes created before processing current node
18 return: List of new nodes created
19 """
20 for _, node_input in enumerate(self.node.input):
21 dequantize_node = self.quantizer._dequantize_value(node_input)
22 if dequantize_node is not None:
23 self.quantizer.new_nodes.append(dequantize_node)
24
25 # Append the original node
26 self.quantizer.new_nodes.append(self.node)
27 