Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_qordered_gelu.py119 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6from logging import getLogger
7
8from fusion_base import Fusion
9from fusion_utils import FusionUtils
10from onnx import helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionQOrderedGelu(Fusion):
17    def __init__(self, model: OnnxModel):
18        super().__init__(model, "QOrderedGelu", ["Gelu", "FastGelu"])
19
20    def fuse(self, node, input_name_to_nodes: dict, output_name_to_node: dict):
21        """
22        INPUT PATTERN
23        Fuse (quantized) Gelu subgraph into one node QOrderedGelu:
24            -> quantized input  -> DQ -> Gelu -> Q ->
25
26        (or)
27
28            -> quantized input  -> DQ -> FastGelu -> Q ->
29
30        OUTPUT PATTERN
31            -> QOrderedGelu ->
32        """
33        gelu_children = self.model.get_children(node, input_name_to_nodes)
34
35        # Should only have 1 child - QuantizeLinear (or)
36        # Should have 2 children - QuantizeLinear + Shape
37        if not (
38            (len(gelu_children) == 1 and gelu_children[0].op_type == "QuantizeLinear")
39            or (
40                len(gelu_children) == 2
41                and gelu_children[0].op_type == "QuantizeLinear"
42                and gelu_children[1].op_type == "Shape"
43            )
44        ):
45            return
46
47        downstream_quantize_node = gelu_children[0]
48        downstream_shape_node = None
49
50        if len(gelu_children) == 2:
51            downstream_shape_node = gelu_children[1]
52
53        if not FusionUtils.check_qdq_node_for_fusion(downstream_quantize_node, self.model):
54            return
55
56        # The first input to Gelu should flow through a DequantizeLinear node
57        first_path_id, first_input_parent_nodes, _ = self.model.match_parent_paths(
58            node,
59            [(["DequantizeLinear"], [0])],
60            output_name_to_node,
61        )
62
63        if first_path_id < 0:
64            return
65
66        upstream_dequantize_node = first_input_parent_nodes[0]
67
68        if not FusionUtils.check_qdq_node_for_fusion(upstream_dequantize_node, self.model):
69            return
70
71        # Fusion logic
72        subgraph_nodes = [node]  # Gelu/FastGelu
73        subgraph_nodes.extend([downstream_quantize_node, upstream_dequantize_node])  # Relevant Q, DQ nodes
74
75        if not self.model.is_safe_to_fuse_nodes(
76            subgraph_nodes,
77            (
78                [node.output[0], downstream_quantize_node.output[0]]
79                if downstream_shape_node is not None
80                else downstream_quantize_node.output
81            ),
82            input_name_to_nodes,
83            output_name_to_node,
84        ):
85            logger.debug("It is not safe to fuse QOrderedGelu node. Skip")
86            return
87
88        self.nodes_to_remove.extend(subgraph_nodes)
89
90        ordered_gelu_node = helper.make_node(
91            "QOrderedGelu",
92            inputs=[
93                upstream_dequantize_node.input[0],
94                upstream_dequantize_node.input[1],
95                downstream_quantize_node.input[1],
96            ],
97            outputs=[downstream_quantize_node.output[0]],
98            name=self.model.create_node_name("QOrderedGelu", name_prefix="QOrderedGelu"),
99        )
100
101        # Arrange the downstream Shape's input to be fed from the
102        # downstream QuantizeLinear node, so that fusion will
103        # be deemed safe
104        if downstream_shape_node is not None:
105            self.model.replace_node_input(
106                downstream_shape_node, downstream_shape_node.input[0], downstream_quantize_node.output[0]
107            )
108
109        # TODO: We only support CuBlasLt order ORDER_ROW for now.
110        # Once we start supporting other data ordering format(s), we
111        # will support user configuring the data ordering for the op.
112        ordered_gelu_node.attribute.extend([helper.make_attribute("order_X", 1)])
113        ordered_gelu_node.attribute.extend([helper.make_attribute("order_Y", 1)])
114
115        ordered_gelu_node.domain = "com.microsoft"
116
117        self.nodes_to_add.append(ordered_gelu_node)
118        self.node_name_to_graph_name[ordered_gelu_node.name] = self.this_graph_name
119 
codekingpro/portable-devtools · Team Ai