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