Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_attention_sam2.py534 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
7import numpy as np
8from fusion_base import Fusion
9from fusion_utils import NumpyHelper
10from onnx import NodeProto, helper, numpy_helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionMultiHeadAttentionSam2(Fusion):
17    """
18    Fuse MultiHeadAttention subgraph of Segment Anything v2 (SAM2).
19    """
20
21    def __init__(
22        self,
23        model: OnnxModel,
24        hidden_size: int,
25        num_heads: int,
26    ):
27        super().__init__(model, "MultiHeadAttention", ["LayerNormalization"])
28        self.hidden_size = hidden_size
29        self.num_heads = num_heads
30
31        # Flags to show warning only once
32        self.num_heads_warning = True
33        self.hidden_size_warning = True
34
35    def get_decoder_num_heads(self, reshape_q: NodeProto) -> int:
36        """Detect num_heads from a reshape node.
37
38        Args:
39            reshape_q (NodeProto): reshape node for Q
40        Returns:
41            int: num_heads, or 0 if not found
42        """
43        num_heads = 0
44
45        # we assume that reshape fusion has done, so the shape is a tensor like [0, 0, num_heads, head_size]
46        shape_value = self.model.get_constant_value(reshape_q.input[1])
47        if shape_value is not None:
48            if isinstance(shape_value, np.ndarray) and list(shape_value.shape) == [4]:
49                num_heads = int(shape_value[2])
50
51        if isinstance(num_heads, int) and num_heads > 0:
52            return num_heads
53
54        return 0
55
56    def get_encoder_num_heads(self, reshape_in: NodeProto) -> int:
57        """Detect num_heads from a reshape node.
58
59        Args:
60            reshape_q (NodeProto): reshape node for Q
61        Returns:
62            int: num_heads, or 0 if not found
63        """
64        num_heads = 0
65
66        shape_value = self.model.get_constant_value(reshape_in.input[1])
67        if shape_value is not None:
68            if isinstance(shape_value, np.ndarray) and list(shape_value.shape) == [5]:
69                num_heads = int(shape_value[3])
70        else:
71            concat_shape = self.model.match_parent(reshape_in, "Concat", 1)
72            if concat_shape is not None and len(concat_shape.input) == 5:
73                # we assume that reshape fusion has done, so the shape is a tensor like [0, 0, num_heads, head_size]
74                shape_value = self.model.get_constant_value(concat_shape.input[3])
75                if shape_value is not None:
76                    if isinstance(shape_value, np.ndarray) and list(shape_value.shape) == [1]:
77                        num_heads = int(shape_value[0])
78
79        if isinstance(num_heads, int) and num_heads > 0:
80            return num_heads
81
82        return 0
83
84    def get_hidden_size(self, layernorm_node):
85        """Detect hidden_size from LayerNormalization node.
86        Args:
87            layernorm_node (NodeProto): LayerNormalization node before Q, K and V
88        Returns:
89            int: hidden_size, or 0 if not found
90        """
91        layernorm_bias = self.model.get_initializer(layernorm_node.input[2])
92        if layernorm_bias:
93            return NumpyHelper.to_array(layernorm_bias).shape[0]
94
95        return 0
96
97    def get_num_heads_and_hidden_size(
98        self, reshape_q: NodeProto, layernorm_node: NodeProto, is_encoder: bool = False
99    ) -> tuple[int, int]:
100        """Detect num_heads and hidden_size.
101
102        Args:
103            reshape_q (NodeProto): reshape node for Q
104            layernorm_node (NodeProto): LayerNormalization node before Q, K, V
105        Returns:
106            Tuple[int, int]: num_heads and hidden_size
107        """
108        if is_encoder:
109            num_heads = self.get_encoder_num_heads(reshape_q)
110        else:
111            num_heads = self.get_decoder_num_heads(reshape_q)
112        if num_heads <= 0:
113            num_heads = self.num_heads  # Fall back to user specified value
114
115        if self.num_heads > 0 and num_heads != self.num_heads:
116            if self.num_heads_warning:
117                logger.warning(f"--num_heads is {self.num_heads}. Detected value is {num_heads}. Using detected value.")
118                self.num_heads_warning = False  # Do not show the warning more than once
119
120        hidden_size = self.get_hidden_size(layernorm_node)
121        if hidden_size <= 0:
122            hidden_size = self.hidden_size  # Fall back to user specified value
123
124        if self.hidden_size > 0 and hidden_size != self.hidden_size:
125            if self.hidden_size_warning:
126                logger.warning(
127                    f"--hidden_size is {self.hidden_size}. Detected value is {hidden_size}. Using detected value."
128                )
129                self.hidden_size_warning = False  # Do not show the warning more than once
130
131        return num_heads, hidden_size
132
133    def create_attention_node(
134        self,
135        q_matmul: NodeProto,
136        q_add: NodeProto,
137        k_matmul: NodeProto,
138        k_add: NodeProto,
139        v_matmul: NodeProto,
140        v_add: NodeProto,
141        num_heads: int,
142        hidden_size: int,
143        output: str,
144    ) -> NodeProto | None:
145        """Create an Attention node.
146
147        Args:
148            q_matmul (NodeProto): MatMul node in fully connection for Q
149            q_add (NodeProto): Add bias node in fully connection for Q
150            k_matmul (NodeProto): MatMul node in fully connection for K
151            k_add (NodeProto): Add bias node in fully connection for K
152            v_matmul (NodeProto): MatMul node in fully connection for V
153            v_add (NodeProto): Add bias node in fully connection for V
154            num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
155            hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
156            output (str): output name
157
158        Returns:
159            Union[NodeProto, None]: the node created or None if failed.
160        """
161        if hidden_size > 0 and (hidden_size % num_heads) != 0:
162            logger.debug(f"input hidden size {hidden_size} is not a multiple of num of heads {num_heads}")
163            return None
164
165        q_weight = self.model.get_initializer(q_matmul.input[1])
166        k_weight = self.model.get_initializer(k_matmul.input[1])
167        v_weight = self.model.get_initializer(v_matmul.input[1])
168        if not (q_weight and k_weight and v_weight):
169            return None
170
171        qw = NumpyHelper.to_array(q_weight)
172        kw = NumpyHelper.to_array(k_weight)
173        vw = NumpyHelper.to_array(v_weight)
174        logger.debug(f"qw={qw.shape} kw={kw.shape} vw={vw.shape} hidden_size={hidden_size}")
175
176        attention_node_name = self.model.create_node_name("MultiHeadAttention")
177
178        attention_inputs = [
179            q_add.output[0],
180            k_add.output[0],
181            v_add.output[0],
182        ]
183
184        attention_node = helper.make_node(
185            "MultiHeadAttention",
186            inputs=attention_inputs,
187            outputs=[output],
188            name=attention_node_name,
189        )
190        attention_node.domain = "com.microsoft"
191        attention_node.attribute.extend([helper.make_attribute("num_heads", num_heads)])
192
193        counter_name = "MultiHeadAttention ({})".format("cross attention")
194        self.increase_counter(counter_name)
195        return attention_node
196
197    def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
198        if self.fuse_sam_encoder_pattern(normalize_node, input_name_to_nodes, output_name_to_node):
199            return
200
201        match_qkv = self.match_attention_subgraph(normalize_node)
202        if match_qkv is None:
203            if normalize_node.input[0] not in output_name_to_node:
204                return
205
206            skip_add = output_name_to_node[normalize_node.input[0]]
207            if skip_add.op_type != "Add":
208                return
209
210            match_qkv = self.match_attention_subgraph(skip_add)
211
212            if match_qkv is None:
213                return
214
215        reshape_qkv, transpose_qkv, reshape_q, matmul_q, add_q, matmul_k, add_k, matmul_v, add_v = match_qkv
216
217        attention_last_node = reshape_qkv
218
219        q_num_heads, q_hidden_size = self.get_num_heads_and_hidden_size(reshape_q, normalize_node, False)
220        if q_num_heads <= 0:
221            logger.debug("fuse_attention: failed to detect num_heads")
222            return
223
224        # number of heads are same for all the paths, hence to create attention node, we pass the q_num_heads
225        new_node = self.create_attention_node(
226            matmul_q,
227            add_q,
228            matmul_k,
229            add_k,
230            matmul_v,
231            add_v,
232            q_num_heads,
233            q_hidden_size,
234            output=attention_last_node.output[0],
235        )
236        if new_node is None:
237            return
238
239        self.nodes_to_add.append(new_node)
240        self.node_name_to_graph_name[new_node.name] = self.this_graph_name
241
242        self.nodes_to_remove.extend([attention_last_node, transpose_qkv])
243
244        # Use prune graph to remove nodes since they are shared by all attention nodes.
245        self.prune_graph = True
246
247    def match_attention_subgraph(self, node_after_output_projection):
248        """Match Q, K and V paths exported by PyTorch 2.*"""
249        qkv_nodes = self.model.match_parent_path(
250            node_after_output_projection,
251            ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
252            [None, None, None, 0, 0],
253        )
254
255        if qkv_nodes is None:
256            return None
257
258        (_, _, reshape_qkv, transpose_qkv, matmul_qkv) = qkv_nodes
259
260        v_nodes = self.model.match_parent_path(matmul_qkv, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None])
261        if v_nodes is None:
262            logger.debug("fuse_attention: failed to match v path")
263            return None
264        (_, _, add_v, matmul_v) = v_nodes
265
266        qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "MatMul"], [0, 0])
267        if qk_nodes is not None:
268            (_softmax_qk, matmul_qk) = qk_nodes
269        else:
270            logger.debug("fuse_attention: failed to match qk path")
271            return None
272
273        q_nodes = self.model.match_parent_path(
274            matmul_qk, ["Mul", "Transpose", "Reshape", "Add", "MatMul"], [0, None, 0, 0, None]
275        )
276        if q_nodes is None:
277            logger.debug("fuse_attention: failed to match q path")
278            return None
279        (mul_q, _transpose_q, reshape_q, add_q, matmul_q) = q_nodes
280
281        k_nodes = self.model.match_parent_path(
282            matmul_qk, ["Mul", "Transpose", "Reshape", "Add", "MatMul"], [1, None, 0, 0, None]
283        )
284        if k_nodes is None:
285            logger.debug("fuse_attention: failed to match k path")
286            return None
287
288        (_mul_k, _, _, add_k, matmul_k) = k_nodes
289
290        # The scalar for Q and K is sqrt(1.0/sqrt(head_size)).
291        mul_q_nodes = self.model.match_parent_path(
292            mul_q,
293            ["Sqrt", "Div", "Sqrt", "Cast", "Slice", "Shape", "Transpose", "Reshape"],
294            [None, 0, 1, 0, 0, 0, 0, 0],
295        )
296        if mul_q_nodes is None or mul_q_nodes[-1] != reshape_q:
297            logger.debug("fuse_attention: failed to match mul_q path")
298            return None
299
300        return reshape_qkv, transpose_qkv, reshape_q, matmul_q, add_q, matmul_k, add_k, matmul_v, add_v
301
302    # --------------------------------------------------------
303    # The following are for SAM encoder
304    # --------------------------------------------------------
305    def fuse_sam_encoder_pattern(self, normalize_node, input_name_to_nodes, output_name_to_node) -> bool:
306        # SAM encoder attention layer pattern:
307        #           Add -----------+
308        #            |             |
309        #        LayerNorm         |
310        #            |             |
311        #        Reshape           |
312        #            |             |
313        #        Transpose         |
314        #            |             |
315        #        MatMul            |
316        #            |             |
317        #           Add            |
318        #            |             |
319        #         Reshape          |
320        #            |             |
321        #          Split           |
322        #            |             |
323        #  Self Attention subgraph |
324        #            |             |
325        #        Reshape           |
326        #            |             |
327        #        Transpose         |
328        #            |             |
329        #        Reshape           |
330        #            |             |
331        #            Add ----------+
332        #            |
333        #         LayerNorm (starts from here)
334
335        nodes = self.model.match_parent_path(
336            normalize_node,
337            ["Add", "Reshape", "Transpose", "Reshape"],
338            [0, None, 0, 0],
339        )
340        if nodes is None:
341            nodes = self.model.match_parent_path(
342                normalize_node,
343                ["Add", "Slice", "Slice", "Reshape", "Transpose", "Reshape"],
344                [0, None, 0, 0, 0, 0],
345            )
346        if nodes is None:
347            nodes = self.model.match_parent_path(
348                normalize_node,
349                ["Add"],
350                [0],
351            )
352        if nodes is None:
353            return False
354
355        node_after_output_projection = nodes[-1]
356        matched_sdpa = self.match_sam_encoder_attention_subgraph(
357            node_after_output_projection, input_index=1 if len(nodes) == 1 else None
358        )
359        if matched_sdpa is None:
360            return False
361
362        reshape_out, transpose_out, split_qkv, transpose_q, transpose_k, transpose_v = matched_sdpa
363
364        # B, S, N, H => B, N, S, H
365        permutation_q = OnnxModel.get_node_attribute(transpose_q, "perm")
366        if (not isinstance(permutation_q, list)) or permutation_q != [0, 2, 1, 3]:
367            return False
368
369        # B, S, N, H => B, N, H, S
370        permutation_k = OnnxModel.get_node_attribute(transpose_k, "perm")
371        if (not isinstance(permutation_k, list)) or permutation_k != [0, 2, 3, 1]:
372            return False
373
374        # B, S, N, H => B, N, S, H
375        permutation_v = OnnxModel.get_node_attribute(transpose_v, "perm")
376        if (not isinstance(permutation_v, list)) or permutation_v != [0, 2, 1, 3]:
377            return False
378
379        input_projection_nodes = self.model.match_parent_path(
380            split_qkv,
381            ["Reshape", "Add", "MatMul"],
382            [0, 0, None],
383        )
384        if input_projection_nodes is None:
385            return False
386        reshape_in, add_in, matmul_in = input_projection_nodes
387        q_num_heads, q_hidden_size = self.get_num_heads_and_hidden_size(reshape_in, normalize_node, True)
388        if q_num_heads <= 0:
389            logger.debug("fuse_attention: failed to detect num_heads")
390            return False
391
392        # Add a shape to convert 4D BxSxNxH to 3D BxSxD, which is required by MHA operator.
393        new_dims_name = "bsnh_to_bsd_reshape_dims"
394        new_dims = self.model.get_initializer(new_dims_name)
395        if new_dims is None:
396            new_dims = numpy_helper.from_array(np.array([0, 0, -1], dtype="int64"), name=new_dims_name)
397            self.model.add_initializer(new_dims, self.this_graph_name)
398        reshape_q_name = self.model.create_node_name("Reshape")
399        reshape_q = helper.make_node(
400            "Reshape",
401            inputs=[transpose_q.input[0], new_dims_name],
402            outputs=[transpose_q.input[0] + "_BSD"],
403            name=reshape_q_name,
404        )
405        self.nodes_to_add.append(reshape_q)
406        self.node_name_to_graph_name[reshape_q.name] = self.this_graph_name
407
408        # Reuse the transpose_q node to transpose K from BSNH to BNSH. Here we update the input and output of the node.
409        transpose_k_bnsh = transpose_q
410        transpose_k_bnsh.input[0] = transpose_k.input[0]
411        transpose_k_bnsh.output[0] = transpose_k.input[0] + "_BNSH"
412
413        logger.debug(f"Found MHA: {q_num_heads=} {q_hidden_size=}")
414
415        # number of heads are same for all the paths, hence to create attention node, we pass the q_num_heads
416        new_node = self.create_mha_node(
417            reshape_q,
418            transpose_k_bnsh,
419            transpose_v,
420            q_num_heads,
421        )
422        if new_node is None:
423            return False
424
425        # Update the input of the next node that consumes the output of the MHA.
426        assert len(self.model.get_children(transpose_out, input_name_to_nodes)) == 1
427        reshape_out.input[0] = new_node.output[0]
428
429        self.nodes_to_add.append(new_node)
430        self.node_name_to_graph_name[new_node.name] = self.this_graph_name
431        self.nodes_to_remove.extend([transpose_out])
432
433        # Use prune graph to remove nodes since they are shared by all attention nodes.
434        self.prune_graph = True
435        return True
436
437    def match_sam_encoder_attention_subgraph(self, node_after_output_projection, input_index=None):
438        """Match SDPA pattern in SAM2 enconder.*"""
439
440        # nodes of output projection and the second MatMul in SDPA.
441        out_nodes = self.model.match_parent_path(
442            node_after_output_projection,
443            ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
444            [input_index, None, None, 0, 0],
445        )
446
447        if out_nodes is None:
448            return None
449
450        (_, _, reshape_out, transpose_out, matmul_qk_v) = out_nodes
451
452        # Split and Reshape is for packed QKV
453        v_nodes = self.model.match_parent_path(matmul_qk_v, ["Transpose", "Squeeze", "Split", "Reshape"], [1, 0, 0, 0])
454        if v_nodes is None:
455            logger.debug("failed to match v path")
456            return None
457        (transpose_v, _, split_qkv, reshape_qkv) = v_nodes
458
459        qk_nodes = self.model.match_parent_path(matmul_qk_v, ["Softmax", "MatMul"], [0, 0])
460        if qk_nodes is not None:
461            (_softmax_qk, matmul_qk) = qk_nodes
462        else:
463            logger.debug("failed to match qk path")
464            return None
465
466        q_nodes = self.model.match_parent_path(matmul_qk, ["Mul", "Transpose", "Squeeze", "Split"], [0, None, 0, 0])
467        if q_nodes is None:
468            q_nodes = self.model.match_parent_path(
469                matmul_qk,
470                ["Mul", "Transpose", "Reshape", "Transpose", "MaxPool", "Transpose", "Reshape", "Squeeze", "Split"],
471                [0, None, 0, 0, 0, 0, 0, 0, 0],
472            )
473            if q_nodes is None:
474                logger.debug("failed to match q path")
475                return None
476
477        if q_nodes[-1] != split_qkv:
478            return None
479        transpose_q = q_nodes[1]
480
481        k_nodes = self.model.match_parent_path(matmul_qk, ["Mul", "Transpose", "Squeeze", "Split"], [1, None, 0, 0])
482        if k_nodes is None:
483            logger.debug("failed to match k path")
484            return None
485
486        if k_nodes[-1] != split_qkv:
487            return None
488        (mul_k, transpose_k, _squeeze_k, _) = k_nodes
489
490        return reshape_out, transpose_out, split_qkv, transpose_q, transpose_k, transpose_v
491
492    def create_mha_node(
493        self,
494        reshape_q: NodeProto,
495        transpose_k: NodeProto,
496        transpose_v: NodeProto,
497        num_heads: int,
498    ) -> NodeProto:
499        """Create a MultiHeadAttention node for SAM2 encoder.
500
501        Args:
502            reshape_q (NodeProto): Reshape node for Q, output is 3D BxSxNH format
503            transpose_k (NodeProto): Transpose node for K, output is BNSH format
504            transpose_v (NodeProto): Transpose node for V, output is BNSH format
505            num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
506
507        Returns:
508            NodeProto: the MultiHeadAttention node created.
509        """
510
511        attention_node_name = self.model.create_node_name("MultiHeadAttention")
512
513        inputs = [
514            reshape_q.output[0],
515            transpose_k.output[0],
516            transpose_v.output[0],
517        ]
518
519        # Create a new output name since the shape is 3D, which is different from the original output shape (4D).
520        output = attention_node_name + "_out"
521
522        attention_node = helper.make_node(
523            "MultiHeadAttention",
524            inputs=inputs,
525            outputs=[output],
526            name=attention_node_name,
527        )
528        attention_node.domain = "com.microsoft"
529        attention_node.attribute.extend([helper.make_attribute("num_heads", num_heads)])
530
531        counter_name = "MultiHeadAttention ({})".format("self attention")
532        self.increase_counter(counter_name)
533        return attention_node
534 
codekingpro/portable-devtools · Team Ai