codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import logging
6
7from fusion_attention import AttentionMask
8from fusion_bart_attention import FusionBartAttention
9from fusion_options import FusionOptions
10from fusion_reshape import FusionReshape
11from onnx import numpy_helper
12from onnx_model import OnnxModel
13from onnx_model_bert import BertOnnxModel
14
15logger = logging.getLogger(__name__)
16
17
18class FusionBartReshape(FusionReshape):
19 def __init__(self, model: OnnxModel):
20 super().__init__(model)
21
22 def fuse(self, reshape_node, input_name_to_nodes, output_name_to_node):
23 if reshape_node.input[1] not in output_name_to_node:
24 return
25
26 concat_node = output_name_to_node[reshape_node.input[1]]
27 if concat_node.op_type != "Concat" or len(concat_node.input) != 4:
28 return
29
30 path0 = self.model.match_parent_path(
31 concat_node,
32 ["Unsqueeze", "Gather", "Shape"],
33 [0, 0, 0],
34 output_name_to_node,
35 )
36 if path0 is None:
37 return
38
39 (_, gather_0, shape_0) = path0
40
41 shape = []
42 gather_value = self.model.get_constant_value(gather_0.input[1])
43 if gather_value == 0:
44 shape.append(0)
45
46 path1 = self.model.match_parent_path(
47 concat_node,
48 ["Unsqueeze", "Gather", "Shape"],
49 [1, 0, 0],
50 output_name_to_node,
51 )
52 if path1 is None:
53 input_1_proto = self.model.get_initializer(concat_node.input[1])
54 input_2_proto = self.model.get_initializer(concat_node.input[2])
55 input_3_proto = self.model.get_initializer(concat_node.input[3])
56 if input_1_proto is None or input_2_proto is None or input_3_proto is None:
57 return
58
59 input_1 = numpy_helper.to_array(input_1_proto)
60 input_2 = numpy_helper.to_array(input_2_proto)
61 input_3 = numpy_helper.to_array(input_3_proto)
62 if len(input_1) != 1 or len(input_2) != 1 or len(input_3) != 1:
63 return
64
65 if not (input_1[0] == -1 and input_2[0] > 0 and input_3[0] > 0):
66 return
67
68 shape.extend(input_1)
69 shape.extend(input_2)
70 shape.extend(input_3)
71 gemm_path_with_bias = self.model.match_parent_path(
72 reshape_node, ["Add", "MatMul"], [0, 1], output_name_to_node
73 )
74 gemm_path_no_bias = self.model.match_parent_path(reshape_node, ["MatMul"], [0], output_name_to_node)
75 if gemm_path_with_bias is not None:
76 gemm_path = gemm_path_with_bias
77 elif gemm_path_no_bias is not None:
78 gemm_path = gemm_path_no_bias
79 else:
80 return
81
82 top_matmul = gemm_path[-1]
83 root_input = top_matmul.input[0]
84
85 self.replace_reshape_node(shape, reshape_node, concat_node)
86 else:
87 (_, gather_1, shape_1) = path1
88
89 gather_value = self.model.get_constant_value(gather_1.input[1])
90 if gather_value == 1:
91 shape.append(0)
92
93 input_2_proto = self.model.get_initializer(concat_node.input[2])
94 input_3_proto = self.model.get_initializer(concat_node.input[3])
95 if input_2_proto is None or input_3_proto is None:
96 return
97
98 input_2 = numpy_helper.to_array(input_2_proto)
99 input_3 = numpy_helper.to_array(input_3_proto)
100 if len(input_2) != 1 or len(input_3) != 1:
101 return
102
103 if not (input_2[0] > 0 and input_3[0] > 0):
104 return
105
106 shape.extend(input_2)
107 shape.extend(input_3)
108 gemm_path = self.model.match_parent_path(
109 reshape_node, ["Mul", "Add", "MatMul"], [0, 0, 1], output_name_to_node
110 )
111 if gemm_path is None:
112 return
113
114 top_matmul = gemm_path[-1]
115 root_input = top_matmul.input[0]
116 if shape_0.input[0] != root_input or shape_1.input[0] != root_input:
117 return
118
119 self.replace_reshape_node(shape, reshape_node, concat_node)
120
121
122class BartOnnxModel(BertOnnxModel):
123 def __init__(self, model, num_heads, hidden_size, model_impl="hf"):
124 super().__init__(model, num_heads, hidden_size)
125 self.attention_mask = AttentionMask(self)
126 self.attention_fusion = FusionBartAttention(self, self.hidden_size, self.num_heads, self.attention_mask)
127 self.bart_reshape_fusion_preprocess = FusionBartReshape(self)
128
129 def optimize(self, options: FusionOptions | None = None, add_dynamic_axes: bool = False):
130 self.attention_fusion.use_multi_head_attention = False if options is None else options.use_multi_head_attention
131 self.attention_fusion.disable_multi_head_attention_bias = (
132 False if options is None else options.disable_multi_head_attention_bias
133 )
134 super().optimize(options, add_dynamic_axes)
135
136 def fuse_attention(self):
137 self.attention_fusion.apply()
138
139 def preprocess(self):
140 self.adjust_reshape_and_expand()
141 self.bart_reshape_fusion_preprocess.apply()
142 