codekingpro/portable-devtools
114k
1from .operators.activation import QDQRemovableActivation, QLinearActivation
2from .operators.argmax import QArgMax
3from .operators.attention import AttentionQuant
4from .operators.base_operator import QuantOperatorBase
5from .operators.binary_op import QLinearBinaryOp
6from .operators.concat import QLinearConcat
7from .operators.conv import ConvInteger, QDQConv, QLinearConv
8from .operators.direct_q8 import Direct8BitOp, QDQDirect8BitOp
9from .operators.embed_layernorm import EmbedLayerNormalizationQuant
10from .operators.gather import GatherQuant, QDQGather
11from .operators.gavgpool import QGlobalAveragePool
12from .operators.gemm import QDQGemm, QLinearGemm
13from .operators.lstm import LSTMQuant
14from .operators.matmul import MatMulInteger, QDQMatMul, QLinearMatMul
15from .operators.maxpool import QDQMaxPool, QMaxPool
16from .operators.norm import QDQNormalization
17from .operators.pad import QDQPad, QPad
18from .operators.pooling import QLinearPool
19from .operators.qdq_base_operator import QDQOperatorBase
20from .operators.resize import QDQResize, QResize
21from .operators.softmax import QLinearSoftmax
22from .operators.split import QDQSplit, QSplit
23from .operators.where import QDQWhere, QLinearWhere
24from .quant_utils import QuantizationMode
25
26CommonOpsRegistry = {
27 "Gather": GatherQuant,
28 "Transpose": Direct8BitOp,
29 "EmbedLayerNormalization": EmbedLayerNormalizationQuant,
30}
31
32IntegerOpsRegistry = {
33 "Conv": ConvInteger,
34 "MatMul": MatMulInteger,
35 "Attention": AttentionQuant,
36 "LSTM": LSTMQuant,
37}
38IntegerOpsRegistry.update(CommonOpsRegistry)
39
40QLinearOpsRegistry = {
41 "ArgMax": QArgMax,
42 "Conv": QLinearConv,
43 "Gemm": QLinearGemm,
44 "MatMul": QLinearMatMul,
45 "Add": QLinearBinaryOp,
46 "Mul": QLinearBinaryOp,
47 "Relu": QLinearActivation,
48 "Clip": QLinearActivation,
49 "LeakyRelu": QLinearActivation,
50 "Sigmoid": QLinearActivation,
51 "MaxPool": QMaxPool,
52 "GlobalAveragePool": QGlobalAveragePool,
53 "Split": QSplit,
54 "Pad": QPad,
55 "Reshape": Direct8BitOp,
56 "Squeeze": Direct8BitOp,
57 "Unsqueeze": Direct8BitOp,
58 "Resize": QResize,
59 "AveragePool": QLinearPool,
60 "Concat": QLinearConcat,
61 "Softmax": QLinearSoftmax,
62 "Where": QLinearWhere,
63}
64QLinearOpsRegistry.update(CommonOpsRegistry)
65
66QDQRegistry = {
67 "Conv": QDQConv,
68 "ConvTranspose": QDQConv,
69 "Gemm": QDQGemm,
70 "Clip": QDQRemovableActivation,
71 "Relu": QDQRemovableActivation,
72 "Reshape": QDQDirect8BitOp,
73 "Transpose": QDQDirect8BitOp,
74 "Squeeze": QDQDirect8BitOp,
75 "Unsqueeze": QDQDirect8BitOp,
76 "Resize": QDQResize,
77 "MaxPool": QDQMaxPool,
78 "AveragePool": QDQDirect8BitOp,
79 "Slice": QDQDirect8BitOp,
80 "Pad": QDQPad,
81 "MatMul": QDQMatMul,
82 "Split": QDQSplit,
83 "Gather": QDQGather,
84 "GatherElements": QDQGather,
85 "Where": QDQWhere,
86 "InstanceNormalization": QDQNormalization,
87 "LayerNormalization": QDQNormalization,
88 "BatchNormalization": QDQNormalization,
89 "TopK": QDQDirect8BitOp,
90 "CumSum": QDQOperatorBase,
91}
92
93
94def CreateDefaultOpQuantizer(onnx_quantizer, node): # noqa: N802
95 return QuantOperatorBase(onnx_quantizer, node)
96
97
98def CreateOpQuantizer(onnx_quantizer, node): # noqa: N802
99 registry = IntegerOpsRegistry if onnx_quantizer.mode == QuantizationMode.IntegerOps else QLinearOpsRegistry
100 if node.op_type in registry:
101 op_quantizer = registry[node.op_type](onnx_quantizer, node)
102 if op_quantizer.should_quantize():
103 return op_quantizer
104 return QuantOperatorBase(onnx_quantizer, node)
105
106
107def CreateQDQQuantizer(onnx_quantizer, node): # noqa: N802
108 if node.op_type in QDQRegistry:
109 return QDQRegistry[node.op_type](onnx_quantizer, node)
110 return QDQOperatorBase(onnx_quantizer, node)
111 