codekingpro/portable-devtools
114k
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 