Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_attention_clip.py341 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from logging import getLogger
6
7from fusion_attention import AttentionMask, FusionAttention
8from fusion_options import AttentionMaskFormat
9from onnx import NodeProto
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionAttentionClip(FusionAttention):
16    """
17    Fuse Attention subgraph of Clip into one Attention node.
18    """
19
20    def __init__(
21        self,
22        model: OnnxModel,
23        hidden_size: int,
24        num_heads: int,
25    ):
26        attention_mask = AttentionMask(model)
27        attention_mask.mask_format = AttentionMaskFormat.NoMask
28
29        super().__init__(
30            model,
31            hidden_size,
32            num_heads,
33            attention_mask,
34            use_multi_head_attention=False,
35            search_op_types=["SkipLayerNormalization"],
36        )
37
38    def get_num_heads_and_hidden_size(self, reshape_q: NodeProto) -> tuple[int, int]:
39        """Detect num_heads and hidden_size for ONNX model from MiDaS
40        Args:
41            reshape_q (NodeProto): reshape node for q
42        Returns:
43            Tuple[int, int]: num_heads and hidden_size
44        """
45        concat = self.model.match_parent(reshape_q, "Concat", 1)
46        if concat is None or len(concat.input) != 4:
47            return self.num_heads, self.hidden_size
48
49        # The shape is a tensor like [?, ?, num_heads, head_size]
50        num_head_value = self.model.get_constant_value(concat.input[2])
51        if num_head_value is None:
52            return self.num_heads, self.hidden_size  # Fall back to user specified value
53
54        if len(num_head_value) != 1 or num_head_value[0] <= 0:
55            return self.num_heads, self.hidden_size  # Fall back to user specified value
56
57        num_heads = num_head_value[0]
58
59        head_size_value = self.model.get_constant_value(concat.input[3])
60        if head_size_value is None:
61            return self.num_heads, self.hidden_size  # Fall back to user specified value
62
63        if len(head_size_value) != 1 or head_size_value[0] <= 0:
64            return self.num_heads, self.hidden_size  # Fall back to user specified value
65
66        head_size = head_size_value[0]
67
68        hidden_size = num_heads * head_size
69
70        if self.num_heads > 0 and num_heads != self.num_heads:
71            if self.num_heads_warning:
72                logger.warning(f"--num_heads is {self.num_heads}. Detected value is {num_heads}. Using detected value.")
73                self.num_heads_warning = False  # Do not show the warning more than once
74
75        if self.hidden_size > 0 and hidden_size != self.hidden_size:
76            if self.hidden_size_warning:
77                logger.warning(
78                    f"--hidden_size is {self.hidden_size}. Detected value is {hidden_size}. Using detected value."
79                )
80                self.hidden_size_warning = False  # Do not show the warning more than once
81
82        return num_heads, hidden_size
83
84    def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
85        skip_input_index = None
86        node_before_layer_norm = None
87        for i in [1, 0]:
88            parent = self.model.match_parent(normalize_node, "SkipLayerNormalization", i)
89            if parent is not None:
90                skip_input_index = i
91                node_before_layer_norm = parent
92
93        root_input = None
94        if node_before_layer_norm is not None:
95            root_input = node_before_layer_norm.output[0]
96        else:
97            # Deal with the first attention after the embedding layer.
98            for i in [0, 1]:
99                node_before_layer_norm = None
100
101                node_before_layer_norm_1 = self.model.match_parent(normalize_node, "Add", i)
102                node_before_layer_norm_2 = self.model.match_parent(normalize_node, "LayerNormalization", i)
103                if node_before_layer_norm_1 is not None:
104                    #           Add -----------+
105                    #            |             |
106                    #        LayerNorm         |
107                    #            |             |
108                    #        LayerNorm         |
109                    #            |             |
110                    #   Attention subgraph     |
111                    #            |             |
112                    #      SkipLayerNorm ------+
113                    node_before_layer_norm = node_before_layer_norm_1
114                elif node_before_layer_norm_2 is not None:
115                    #           Add
116                    #            |
117                    #        LayerNorm --------+
118                    #            |             |
119                    #        LayerNorm         |
120                    #            |             |
121                    #   Attention subgraph     |
122                    #            |             |
123                    #      SkipLayerNorm ------+
124                    node_before_layer_norm = node_before_layer_norm_2
125
126                if node_before_layer_norm is None:
127                    continue
128                child = self.model.find_first_child_by_type(
129                    node_before_layer_norm,
130                    "LayerNormalization",
131                    input_name_to_nodes,
132                    False,
133                )
134                if child is None:
135                    continue
136                root_input = child.output[0]
137                skip_input_index = i
138                break
139
140            if skip_input_index is None:
141                return
142
143        qkv_nodes = self.model.match_parent_path(
144            normalize_node,
145            ["Add", "MatMul", "Reshape", "Transpose", "Reshape", "MatMul"],
146            [1 - skip_input_index, None, None, 0, 0, 0],
147        )
148        if qkv_nodes is None:
149            qkv_nodes = self.model.match_parent_path(
150                normalize_node,
151                ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
152                [1, None, 0, 0, 0],
153            )
154            if qkv_nodes is None:
155                logger.debug("fuse_attention: failed to match qkv path")
156                return
157        reshape_qkv, transpose_qkv, matmul_qkv = (
158            qkv_nodes[2],
159            qkv_nodes[3],
160            qkv_nodes[-1],
161        )
162
163        v_nodes = self.model.match_parent_path(
164            matmul_qkv,
165            ["Reshape", "Transpose", "Reshape", "Add", "MatMul"],
166            [1, 0, 0, 0, None],
167        )
168        if v_nodes is None:
169            v_nodes = self.model.match_parent_path(
170                matmul_qkv, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None]
171            )
172            if v_nodes is None:
173                logger.debug("fuse_attention: failed to match v path")
174                return
175
176        add_v, matmul_v = v_nodes[-2], v_nodes[-1]
177
178        causal_mask_input_index = None
179        add_mask = None
180        add_mask_indices = []
181        qk_nodes = self.model.match_parent_path(
182            matmul_qkv,
183            ["Softmax", "Reshape", "Add", "Reshape", "MatMul"],
184            [0, 0, 0, None, 0],
185            return_indice=add_mask_indices,
186        )
187        if qk_nodes is None:
188            qk_nodes = self.model.match_parent_path(
189                matmul_qkv,
190                ["Softmax", "MatMul"],
191                [0, 0],
192            )
193            if qk_nodes is None:
194                qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Add", "Mul", "MatMul"], [0, 0, 0, 0])
195                if qk_nodes is not None:
196                    add_mask = qk_nodes[1]
197                else:
198                    # If attention mask is not used, we can still match the qk path.
199                    qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Mul", "MatMul"], [0, 0, 0])
200                    if qk_nodes is None:
201                        # Cast nodes are added in the model for fp16.
202                        qk_nodes = self.model.match_parent_path(
203                            matmul_qkv,
204                            ["Cast", "Cast", "Softmax", "Add", "Mul", "MatMul"],
205                            [0, 0, 0, 0, 0, 0],
206                        )
207                        if qk_nodes is not None:
208                            add_mask = qk_nodes[3]
209                        else:
210                            # If attention mask is not used, we can still match the qk path.
211                            qk_nodes = self.model.match_parent_path(
212                                matmul_qkv,
213                                ["Cast", "Cast", "Softmax", "Mul", "MatMul"],
214                                [0, 0, 0, 0, 0],
215                            )
216                            if qk_nodes is None:
217                                logger.debug("fuse_attention: failed to match qk path")
218                                return
219        else:
220            assert len(add_mask_indices) == 1
221            causal_mask_input_index = 1 - add_mask_indices[0]
222            add_mask = qk_nodes[2]
223
224        matmul_qk = qk_nodes[-1]
225
226        q_nodes = self.model.match_parent_path(
227            matmul_qk,
228            ["Reshape", "Transpose", "Reshape", "Mul", "Add", "MatMul"],
229            [0, 0, 0, 0, None, None],
230        )
231        if q_nodes is None:
232            q_nodes = self.model.match_parent_path(
233                matmul_qk, ["Transpose", "Reshape", "Add", "MatMul"], [0, 0, 0, None]
234            )
235            if q_nodes is None:
236                logger.debug("fuse_attention: failed to match q path")
237                return
238
239            reshape_q = q_nodes[1]
240        else:
241            reshape_q = q_nodes[2]
242
243        add_q, matmul_q = q_nodes[-2], q_nodes[-1]
244
245        k_nodes = self.model.match_parent_path(
246            matmul_qk,
247            ["Transpose", "Reshape", "Transpose", "Reshape", "Add", "MatMul"],
248            [1, 0, 0, 0, 0, None],
249        )
250        if k_nodes is None:
251            k_nodes = self.model.match_parent_path(
252                matmul_qk, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None]
253            )
254            if k_nodes is None:
255                logger.debug("fuse_attention: failed to match k path")
256                return
257
258        add_k, matmul_k = k_nodes[-2], k_nodes[-1]
259
260        if matmul_q.input[0] != root_input or matmul_k.input[0] != root_input or matmul_v.input[0] != root_input:
261            logger.debug("fuse_attention: expect to have same input to q, k and v matmul")
262            return
263
264        num_heads, hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
265        if num_heads <= 0 or hidden_size <= 0:
266            logger.debug("fuse_attention: failed to detect num_heads or hidden_size")
267            return
268
269        attention_last_node = reshape_qkv
270
271        add_qk = ""
272        causal_mask_nodes_1 = None
273        causal_mask_nodes_2 = None
274        if add_mask is not None:
275            if add_mask.input[1] == "attention_mask":
276                add_qk = add_mask.input[1]
277            else:
278                # 4D Add after Q x K'
279                add_qk_nodes = self.model.match_parent_path(
280                    add_mask,
281                    [
282                        "Where",
283                        "Sub",
284                        "Cast",
285                        "Expand",
286                        "Unsqueeze",
287                        "Unsqueeze",
288                        "Reshape",
289                        "Reshape",
290                        "Cast",
291                    ],
292                    [1, 2, 1, 0, 0, 0, 0, 0, 0],
293                )
294                if add_qk_nodes is not None:
295                    add_qk = add_mask.input[1]
296                else:
297                    # Here we do not match the whole subgraph since it is very complex. Instead, we just check whether a key path
298                    # of computing causal mask.
299                    causal_mask_nodes_1 = self.model.match_parent_path(
300                        add_mask,
301                        ["Concat", "Expand", "Unsqueeze", "Unsqueeze", "Where", "Less"],
302                        [causal_mask_input_index, 0, 0, 0, 0, 0],
303                    )
304                    # If the model is exported with batch_size == 1, there is no Concat node
305                    causal_mask_nodes_2 = self.model.match_parent_path(
306                        add_mask,
307                        ["Expand", "Unsqueeze", "Unsqueeze", "Where", "Less"],
308                        [causal_mask_input_index, 0, 0, 0, 0],
309                    )
310
311                    if causal_mask_nodes_1 is None and causal_mask_nodes_2 is None:
312                        logger.debug("fuse_attention: failed to match causal mask subgraph")
313                        return
314
315        new_node = self.create_attention_node(
316            mask_index=None,
317            q_matmul=matmul_q,
318            k_matmul=matmul_k,
319            v_matmul=matmul_v,
320            q_add=add_q,
321            k_add=add_k,
322            v_add=add_v,
323            num_heads=num_heads,
324            hidden_size=hidden_size,
325            first_input=root_input,
326            output=attention_last_node.output[0],
327            add_qk_str=add_qk,
328            scale=None,
329            causal=(causal_mask_nodes_1 is not None) or (causal_mask_nodes_2 is not None),
330        )
331        if new_node is None:
332            logger.debug("fuse_attention: failed to create fused node")
333            return
334
335        self.nodes_to_add.append(new_node)
336        self.node_name_to_graph_name[new_node.name] = self.this_graph_name
337        self.nodes_to_remove.extend([attention_last_node, transpose_qkv])
338
339        # Use prune graph to remove nodes since they are shared by all attention nodes.
340        self.prune_graph = True
341 
codekingpro/portable-devtools · Team Ai