Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model_bart.py142 linesDownload Raw Back to transformers
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 
codekingpro/portable-devtools · Team Ai