Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
attention.py74 linesDownload Raw Back to operators
1import onnx
2from onnx import onnx_pb as onnx_proto  # noqa: F401
3
4from ..quant_utils import attribute_to_kwarg, ms_domain
5from .base_operator import QuantOperatorBase
6
7"""
8    Quantize Attention
9"""
10
11
12class AttentionQuant(QuantOperatorBase):
13    def __init__(self, onnx_quantizer, onnx_node):
14        super().__init__(onnx_quantizer, onnx_node)
15
16    def should_quantize(self):
17        return self.quantizer.should_quantize_node(self.node)
18
19    def quantize(self):
20        """
21        parameter node: Attention node.
22        parameter new_nodes_list: List of new nodes created before processing this node.
23        return: a list of nodes in topological order that represents quantized Attention node.
24        """
25        node = self.node
26        assert node.op_type == "Attention"
27
28        # TODO This is a temporary fix to stop exporting QAttention with qkv_hidden_sizes
29        # attribute. This needs to be removed once the QAttention for varied q,k,v sizes
30        # is implemented
31        for attr in node.attribute:
32            if attr.name == "qkv_hidden_sizes":
33                return super().quantize()
34
35        (
36            quantized_input_names,
37            zero_point_names,
38            scale_names,
39            nodes,
40        ) = self.quantizer.quantize_activation(node, [0])
41
42        (
43            quantized_input_names_weight,
44            zero_point_names_weight,
45            scale_names_weight,
46            nodes_weight,
47        ) = self.quantizer.quantize_weight(node, [1], reduce_range=True, op_level_per_channel=True)
48        quantized_input_names.extend(quantized_input_names_weight)
49        zero_point_names.extend(zero_point_names_weight)
50        scale_names.extend(scale_names_weight)
51        nodes.extend(nodes_weight)
52
53        if quantized_input_names is None:
54            return super().quantize()
55
56        qattention_name = "" if not node.name else node.name + "_quant"
57
58        inputs = []
59        inputs.extend(quantized_input_names)
60        inputs.extend([node.input[2]])
61        inputs.extend(scale_names)
62        inputs.extend([node.input[3] if len(node.input) > 3 else ""])
63        inputs.extend(zero_point_names)
64        inputs.extend([node.input[4] if len(node.input) > 4 else ""])
65
66        kwargs = {}
67        for attribute in node.attribute:
68            kwargs.update(attribute_to_kwarg(attribute))
69        kwargs["domain"] = ms_domain
70        qattention_node = onnx.helper.make_node("QAttention", inputs, node.output, qattention_name, **kwargs)
71        nodes.append(qattention_node)
72
73        self.quantizer.new_nodes += nodes
74 
codekingpro/portable-devtools · Team Ai