Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_bart_attention.py507 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import logging
6
7import numpy as np
8from fusion_attention import AttentionMask, FusionAttention
9from onnx import helper
10from onnx_model import OnnxModel
11
12logger = logging.getLogger(__name__)
13
14
15class FusionBartAttention(FusionAttention):
16    """
17    Fuse Bart Attention subgraph into one Attention node.
18    """
19
20    def __init__(
21        self,
22        model: OnnxModel,
23        hidden_size: int,
24        num_heads: int,
25        attention_mask: AttentionMask,
26    ):
27        super().__init__(model, hidden_size, num_heads, attention_mask)
28
29    def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
30        # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
31        qkv_nodes = self.model.match_parent_path(
32            normalize_node,
33            ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
34            [1, 1, 0, 0, 0],
35        )
36
37        # For LayerNormalization (when SkipLayerNorm fusion doesn't run, e.g. SDPA models where
38        # symbolic shape inference fails), there's an extra Add node for the residual connection
39        # between the LayerNorm and the attention output path.
40        add_before_layernorm = None
41        if qkv_nodes is None:
42            qkv_nodes_with_residual = self.model.match_parent_path(
43                normalize_node,
44                ["Add", "Add", "MatMul", "Reshape", "Transpose", "MatMul"],
45                [0, None, 0, 0, 0, 0],
46            )
47            if qkv_nodes_with_residual is not None:
48                add_before_layernorm = qkv_nodes_with_residual[0]
49                qkv_nodes = qkv_nodes_with_residual[1:]
50
51        if qkv_nodes is not None:
52            (
53                add_out,
54                matmul_out,
55                reshape_qkv,
56                transpose_qkv,
57                matmul_qkv,
58            ) = qkv_nodes
59        else:
60            logger.debug("fuse_attention: failed to match qkv path")
61            return
62
63        if add_before_layernorm is not None:
64            # LayerNorm case: root_input is the non-attention input of the residual Add
65            if add_before_layernorm.input[0] == add_out.output[0]:
66                root_input = add_before_layernorm.input[1]
67            else:
68                root_input = add_before_layernorm.input[0]
69        else:
70            other_inputs = []
71            for input_ in normalize_node.input:
72                if input_ not in output_name_to_node:
73                    continue
74                if input_ == qkv_nodes[0].output[0]:
75                    continue
76                other_inputs.append(input_)
77            if len(other_inputs) != 1:
78                return
79            root_input = other_inputs[0]
80
81        # Sometimes the input name to the attention MatMul nodes does not match the input name to the end
82        # SkipLayerNormalization node (name saved in root_input). We find the true input name to the MatMul
83        # nodes by getting the initial SkipLayerNormalization node and checking how many MatMul nodes are
84        # children nodes for each of its output names.
85        """
86                                        root_input
87                    +---------------------------------------------------+
88                    |                                                   |
89                    |                                                   |
90        SkipLayerNormalization --> Attention --> MatMul --> SkipLayerNormalization
91        """
92        skip_layernorm = output_name_to_node[root_input]
93        # For some attention blocks, the end SkipLayerNormalization node may point to another node whose
94        # child is the LayerNormalization node.
95        if skip_layernorm.op_type in {"Add", "Clip"}:
96            skip_layernorm = self.model.get_children(skip_layernorm)[0]
97        for output in skip_layernorm.output:
98            if not output:
99                continue
100            children = input_name_to_nodes[output]
101            children_types = [child.op_type for child in children]
102            if children_types.count("MatMul") >= 1:
103                root_input = output
104                break
105
106        graph_input_names = {node.name for node in self.model.graph().input}
107        graph_output_names = {node.name for node in self.model.graph().output}
108
109        v_nodes_past_or_present = self.model.match_parent_path(
110            matmul_qkv,
111            ["Transpose", "Reshape", "Add", "MatMul"],
112            [1, 0, 0, None],
113        )
114        v_nodes_with_past = self.model.match_parent_path(
115            matmul_qkv,
116            ["Concat", "Transpose", "Reshape", "Add", "MatMul"],
117            [1, 1, 0, 0, None],
118        )
119        v_nodes_past_only_oai = self.model.match_parent_path(
120            matmul_qkv,
121            ["Transpose", "Reshape", "Reshape", "Transpose"],
122            [1, 0, 0, 0],
123        )
124        past_v, present_v = "", ""
125        v_nodes, add_v, matmul_v = [], None, None
126        if v_nodes_past_or_present is not None:
127            v_nodes = v_nodes_past_or_present
128            (transpose_v, reshape_v, add_v, matmul_v) = v_nodes
129
130            # Find past_v input name
131            start_child_nodes = input_name_to_nodes[add_v.output[0]]
132            for start_child_node in start_child_nodes:
133                if start_child_node.op_type == "Concat":
134                    concat_v_nodes = self.model.match_parent_path(
135                        start_child_node,
136                        ["Reshape", "Transpose"],
137                        [0, 0],
138                    )
139                    if concat_v_nodes is not None:
140                        past_v = concat_v_nodes[-1].input[0]
141                    start_child_nodes = input_name_to_nodes[start_child_node.output[0]]
142                    break
143
144            # Find present_v output name
145            for start_child_node in start_child_nodes:
146                start_grandchild_nodes = input_name_to_nodes[start_child_node.output[0]]
147                for start_grandchild_node in start_grandchild_nodes:
148                    if start_grandchild_node.output[0] in graph_output_names:
149                        present_v = start_grandchild_node.output[0]
150                        break
151                if present_v != "":
152                    break
153        elif v_nodes_with_past is not None:
154            v_nodes = v_nodes_with_past
155            (concat_v, transpose_v, reshape_v, add_v, matmul_v) = v_nodes
156            past_v = concat_v.input[0]
157            present_v = concat_v.output[0]
158        elif matmul_qkv.input[1] in graph_input_names:
159            # Hugging Face's cross-attention where past_v is used directly as value
160            past_v = matmul_qkv.input[1]
161        elif v_nodes_past_only_oai is not None:
162            # OpenAI's cross-attention where past_v is used directly as value
163            v_nodes = v_nodes_past_only_oai
164            past_v = v_nodes[-1].input[0]
165        else:
166            logger.debug("fuse_attention: failed to match v path")
167            return
168        past_v = past_v if past_v in graph_input_names else ""
169        present_v = present_v if present_v in graph_output_names else ""
170
171        qk_nodes_no_mask = self.model.match_parent_path(matmul_qkv, ["Softmax", "MatMul"], [0, 0])
172        qk_nodes_with_mask = self.model.match_parent_path(matmul_qkv, ["Softmax", "Add", "MatMul"], [0, 0, 0])
173        # SDPA: NaN guard (Where(IsNaN, 0, softmax)) wraps the Softmax output.
174        # Where input[2] is the Softmax output (value when condition is False).
175        qk_nodes_sdpa_no_mask = self.model.match_parent_path(matmul_qkv, ["Where", "Softmax", "MatMul"], [0, 2, 0])
176        qk_nodes_sdpa_with_mask = self.model.match_parent_path(
177            matmul_qkv, ["Where", "Softmax", "Add", "MatMul"], [0, 2, 0, 0]
178        )
179        qk_nodes, add_qk = [], None
180        if qk_nodes_no_mask is not None:
181            _, matmul_qk = qk_nodes_no_mask
182            qk_nodes = qk_nodes_no_mask
183        elif qk_nodes_with_mask is not None:
184            _, add_qk, matmul_qk = qk_nodes_with_mask
185            qk_nodes = qk_nodes_with_mask
186        elif qk_nodes_sdpa_no_mask is not None:
187            _, _, matmul_qk = qk_nodes_sdpa_no_mask
188            qk_nodes = qk_nodes_sdpa_no_mask
189        elif qk_nodes_sdpa_with_mask is not None:
190            _, _, add_qk, matmul_qk = qk_nodes_sdpa_with_mask
191            qk_nodes = qk_nodes_sdpa_with_mask
192        else:
193            logger.debug("fuse_attention: failed to match qk path")
194            return
195
196        q_nodes_hf = self.model.match_parent_path(
197            matmul_qk,
198            ["Transpose", "Reshape", "Mul", "Add", "MatMul"],
199            [0, 0, 0, 0, 1],
200        )
201        q_nodes_oai = self.model.match_parent_path(
202            matmul_qk,
203            ["Mul", "Transpose", "Reshape", "Add", "MatMul"],
204            [0, 0, 0, 0, 1],
205        )
206        # SDPA: Mul(scale) applied before Transpose, MatMul may be at any Add input.
207        q_nodes_sdpa = self.model.match_parent_path(
208            matmul_qk,
209            ["Mul", "Transpose", "Reshape", "Add", "MatMul"],
210            [0, 0, 0, 0, None],
211        )
212        q_nodes = []
213        if q_nodes_hf is not None:
214            q_nodes = q_nodes_hf
215            (transpose_q, reshape_q, mul_q, add_q, matmul_q) = q_nodes
216        elif q_nodes_oai is not None:
217            q_nodes = q_nodes_oai
218            (mul_q, transpose_q, reshape_q, add_q, matmul_q) = q_nodes
219        elif q_nodes_sdpa is not None:
220            q_nodes = q_nodes_sdpa
221            (mul_q, transpose_q, reshape_q, add_q, matmul_q) = q_nodes
222        else:
223            logger.debug("fuse_attention: failed to match q path")
224            return
225
226        k_nodes_no_past_hf = self.model.match_parent_path(
227            matmul_qk,
228            ["Transpose", "Reshape", "MatMul"],
229            [1, 0, 0],
230        )
231        k_nodes_with_past_hf = self.model.match_parent_path(
232            matmul_qk,
233            ["Transpose", "Concat", "Transpose", "Reshape", "MatMul"],
234            [1, 0, 1, 0, 0],
235        )
236        k_nodes_past_or_present_oai = self.model.match_parent_path(
237            matmul_qk,
238            ["Mul", "Transpose", "Reshape", "MatMul"],
239            [1, 0, 0, 0],
240        )
241        k_nodes_past_only_oai = self.model.match_parent_path(
242            matmul_qk,
243            ["Mul", "Transpose", "Reshape", "Reshape", "Transpose"],
244            [1, 0, 0, 0, 0],
245        )
246        # SDPA: K is scaled (Mul) and transposed via Reshape->Transpose(0,2,1)->Reshape chain.
247        k_nodes_sdpa = self.model.match_parent_path(
248            matmul_qk,
249            ["Mul", "Reshape", "Transpose", "Reshape", "Transpose", "Reshape", "Add", "MatMul"],
250            [1, 0, 0, 0, 0, 0, 0, None],
251        )
252        past_k, present_k = "", ""
253        k_nodes, add_k, matmul_k = [], None, None
254        if k_nodes_no_past_hf is not None:
255            k_nodes = k_nodes_no_past_hf
256            (transpose_k, reshape_k, matmul_k) = k_nodes
257
258            # Find present_k output name
259            transpose_k_nodes = input_name_to_nodes[reshape_k.output[0]]
260            for transpose_k_node in transpose_k_nodes:
261                if transpose_k_node.output[0] in graph_output_names:
262                    present_k = transpose_k_node.output[0]
263                    break
264        elif k_nodes_with_past_hf is not None:
265            k_nodes = k_nodes_with_past_hf
266            (_, concat_k, transpose_k, reshape_k, matmul_k) = k_nodes
267            past_k = concat_k.input[0]
268            present_k = concat_k.output[0]
269        elif output_name_to_node[matmul_qk.input[1]].input[0] in graph_input_names:
270            # Hugging Face's cross-attention where past_k is used directly as key
271            k_nodes = [output_name_to_node[matmul_qk.input[1]]]
272            past_k = k_nodes[0].input[0]
273        elif k_nodes_sdpa is not None:
274            k_nodes = k_nodes_sdpa
275            (_, _, _, _, transpose_k, reshape_k, add_k, matmul_k) = k_nodes
276        elif k_nodes_past_or_present_oai is not None:
277            k_nodes = k_nodes_past_or_present_oai
278            (_, transpose_k, reshape_k, matmul_k) = k_nodes
279
280            # Find past_k input name
281            start_child_nodes = input_name_to_nodes[matmul_k.output[0]]
282            for start_child_node in start_child_nodes:
283                if start_child_node.op_type == "Concat":
284                    concat_k_nodes = self.model.match_parent_path(
285                        start_child_node,
286                        ["Reshape", "Transpose"],
287                        [0, 0],
288                    )
289                    if concat_k_nodes is not None:
290                        past_k = concat_k_nodes[-1].input[0]
291                    start_child_nodes = input_name_to_nodes[start_child_node.output[0]]
292                    break
293
294            # Find present_k output name
295            for start_child_node in start_child_nodes:
296                start_grandchild_nodes = input_name_to_nodes[start_child_node.output[0]]
297                for start_grandchild_node in start_grandchild_nodes:
298                    if start_grandchild_node.output[0] in graph_output_names:
299                        present_k = start_grandchild_node.output[0]
300                        break
301                if present_k != "":
302                    break
303        elif k_nodes_past_only_oai is not None:
304            # OpenAI's cross-attention where past_k is used directly as key
305            k_nodes = k_nodes_past_only_oai
306            past_k = k_nodes[-1].input[0]
307        else:
308            logger.debug("fuse_attention: failed to match k path")
309            return
310        past_k = past_k if past_k in graph_input_names else ""
311        present_k = present_k if present_k in graph_output_names else ""
312
313        if matmul_k is not None and add_k is None:
314            # Create empty Add node for attention graph
315            add_v_tensor = self.model.get_initializer(add_v.input[0])
316            bias_dim = add_v_tensor.dims[0]
317            dtype = add_v_tensor.data_type
318            empty_bias_name = "empty_bias"
319            empty_tensor = self.model.get_initializer(empty_bias_name)
320            if empty_tensor is None:
321                self.add_initializer(
322                    empty_bias_name,
323                    dtype,
324                    dims=[bias_dim],
325                    vals=np.array([0.0] * bias_dim, dtype=helper.tensor_dtype_to_np_dtype(dtype)),
326                )
327
328            add_name = self.model.create_node_name("Add")
329            add_k = helper.make_node("Add", [empty_bias_name, matmul_k.output[0]], [reshape_k.name], add_name)
330
331        three_root_inputs = bool(past_k) and bool(past_v) and matmul_k is None and matmul_v is None
332        one_root_input = (
333            not three_root_inputs
334            and matmul_q.input[0] == root_input
335            and matmul_k.input[0] == root_input
336            and matmul_v.input[0] == root_input
337        )
338        two_root_inputs = (
339            not three_root_inputs
340            and matmul_q.input[0] == root_input
341            and matmul_k.input[0] == matmul_v.input[0]
342            and matmul_k.input[0] != matmul_q.input[0]
343        )
344
345        # There are 5 types of attention:
346        # 1) Encoder attention with one_root_input=True and no mask
347        # 2) Decoder self attention with one_root_input=True and has mask
348        # 3) Decoder cross attention with two_root_inputs=True and no mask
349        # 4) Decoder self attention with past with one_root_input=True and has mask and past_k and past_v
350        # 5) Decoder cross attention with past with three_root_inputs=True and no mask
351        # Derive mask presence from which QK pattern matched rather than re-walking the graph.
352        # This reuses the result of match_parent_paths above, which already tried both masked and
353        # unmasked variants and returned the first successful match.
354        has_mask = qk_nodes in (qk_nodes_with_mask, qk_nodes_sdpa_with_mask)
355        no_mask = not has_mask
356        encoder_attention = one_root_input and no_mask
357        decoder_self_attention = one_root_input and has_mask
358        decoder_cross_attention = two_root_inputs and no_mask
359        decoder_self_attention_with_past = decoder_self_attention and bool(past_k) and bool(past_v)
360        decoder_cross_attention_with_past = three_root_inputs and no_mask
361
362        # For decoder self-attentions, the attention mask needs to be included in the attention node
363        causal_mask = has_mask
364        mask_nodes = []
365        if causal_mask:
366            mask_nodes_bart = self.model.match_parent_path(
367                add_qk,
368                ["Where"],
369                [1],
370            )
371            mask_nodes_whisper_hf = self.model.match_parent_path(
372                add_qk,
373                ["Slice", "Expand", "Where"],
374                [1, 0, 1],
375            )
376            mask_nodes_whisper_oai = self.model.match_parent_path(
377                add_qk,
378                ["Slice", "Unsqueeze", "Gather", "Shape", "Add"],
379                [1, 2, 0, 0, 0],
380            )
381            mask_nodes_whisper_oai_unit_test = self.model.match_parent_path(
382                add_qk,
383                ["Slice", "Slice"],
384                [1, 0],
385            )
386            if mask_nodes_whisper_hf is not None:
387                mask_nodes = mask_nodes_whisper_hf
388            elif mask_nodes_whisper_oai is not None:
389                mask_nodes = mask_nodes_whisper_oai
390            elif mask_nodes_whisper_oai_unit_test is not None:
391                mask_nodes = mask_nodes_whisper_oai_unit_test
392            elif mask_nodes_bart is not None:
393                mask_nodes = mask_nodes_bart
394            else:
395                logger.debug("fuse_attention: failed to match mask nodes")
396                return
397            assert len(mask_nodes) > 0
398
399        if (
400            encoder_attention
401            or decoder_self_attention
402            or decoder_cross_attention
403            or decoder_self_attention_with_past
404            or decoder_cross_attention_with_past
405        ):
406            attention_last_node = reshape_qkv
407            num_heads, hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
408
409            # Fall back to user-specified values when detected values are invalid
410            # (e.g., SDPA models use -1 in reshape shapes for dynamic dimensions).
411            if (num_heads <= 0 or hidden_size <= 0) and self.num_heads > 0 and self.hidden_size > 0:
412                logger.debug(
413                    "fuse_attention: reshape dims invalid (num_heads=%d, hidden_size=%d), "
414                    "falling back to user-specified num_heads=%d, hidden_size=%d",
415                    num_heads,
416                    hidden_size,
417                    self.num_heads,
418                    self.hidden_size,
419                )
420                num_heads = self.num_heads
421                hidden_size = self.hidden_size
422
423            if num_heads <= 0 or hidden_size <= 0 or (hidden_size % num_heads) != 0:
424                logger.debug("fuse_attention: failed to detect num_heads or hidden_size")
425                return
426
427            new_node = None
428            if decoder_self_attention_with_past or decoder_cross_attention or decoder_cross_attention_with_past:
429                # Note: Decoder attention with past key and past value is fused as multi-head attention
430                # rather than attention because multi-head attention supports separate past key and past
431                # value whereas attention supports concatenated past key and past value.
432                new_node = (
433                    self.create_multihead_attention_node(
434                        q_matmul=matmul_q,
435                        k_matmul=matmul_k if decoder_cross_attention or decoder_self_attention_with_past else past_k,
436                        v_matmul=matmul_v if decoder_cross_attention or decoder_self_attention_with_past else past_v,
437                        q_add=add_q,
438                        k_add=add_k if decoder_cross_attention or decoder_self_attention_with_past else None,
439                        v_add=add_v if decoder_cross_attention or decoder_self_attention_with_past else None,
440                        num_heads=num_heads,
441                        hidden_size=hidden_size,
442                        output=attention_last_node.output[0],
443                        unidirectional=causal_mask,
444                        past_k=past_k if decoder_self_attention_with_past else "",
445                        past_v=past_v if decoder_self_attention_with_past else "",
446                        present_k=present_k,
447                        present_v=present_v,
448                    )
449                    if self.use_multi_head_attention
450                    else None
451                )
452            else:
453                # Temporarily set multi-head attention flag to false
454                use_multi_head_attention_ground_truth = self.use_multi_head_attention
455                self.use_multi_head_attention = False
456                new_node = self.create_attention_node(
457                    mask_index=None,
458                    q_matmul=matmul_q,
459                    k_matmul=matmul_k,
460                    v_matmul=matmul_v,
461                    q_add=add_q,
462                    k_add=add_k,
463                    v_add=add_v,
464                    num_heads=num_heads,
465                    hidden_size=hidden_size,
466                    first_input=root_input,
467                    output=attention_last_node.output[0],
468                    causal=causal_mask,
469                    past_k=past_k,
470                    past_v=past_v,
471                    present_k=present_k,
472                    present_v=present_v,
473                )
474                self.use_multi_head_attention = use_multi_head_attention_ground_truth
475            if new_node is None:
476                logger.debug("fuse_attention: failed to create fused node")
477                return
478
479            self.nodes_to_add.append(new_node)
480            self.node_name_to_graph_name[new_node.name] = self.this_graph_name
481
482            self.nodes_to_remove.extend([attention_last_node, transpose_qkv, matmul_qkv])
483            self.nodes_to_remove.extend(qk_nodes)
484
485            # When using multi-head attention, keep MatMul nodes in original graph
486            if decoder_self_attention_with_past or decoder_cross_attention or decoder_cross_attention_with_past:
487                if len(q_nodes) > 0 and q_nodes[-1].op_type == "MatMul":
488                    q_nodes.pop()
489                if len(k_nodes) > 0 and k_nodes[-1].op_type == "MatMul":
490                    k_nodes.pop()
491                if len(v_nodes) > 0 and v_nodes[-1].op_type == "MatMul":
492                    v_nodes.pop()
493                if self.disable_multi_head_attention_bias:
494                    if len(q_nodes) > 0 and q_nodes[-1].op_type == "Add":
495                        q_nodes.pop()
496                    if len(k_nodes) > 0 and k_nodes[-1].op_type == "Add":
497                        k_nodes.pop()
498                    if len(v_nodes) > 0 and v_nodes[-1].op_type == "Add":
499                        v_nodes.pop()
500
501            self.nodes_to_remove.extend(q_nodes)
502            self.nodes_to_remove.extend(k_nodes)
503            self.nodes_to_remove.extend(v_nodes)
504
505            # Use prune graph to remove mask nodes since they are shared by all attention nodes.
506            self.prune_graph = True
507 
codekingpro/portable-devtools · Team Ai