Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_attention.py1199 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_options import AttentionMaskFormat
10from fusion_utils import FusionUtils, NumpyHelper
11from onnx import NodeProto, TensorProto, helper, numpy_helper
12from onnx_model import OnnxModel
13
14logger = getLogger(__name__)
15
16
17class AttentionMask:
18    """
19    Fuse Attention subgraph into one Attention node.
20    """
21
22    def __init__(self, model: OnnxModel):
23        self.model = model
24        # A lookup table with mask input as key, and mask index output as value
25        self.mask_indice = {}
26        # A lookup table with mask input as key, and cast (to int32) output as value
27        self.mask_casted = {}
28        self.utils = FusionUtils(model)
29        self.mask_format = AttentionMaskFormat.MaskIndexEnd
30        self.opset_version = model.get_opset_version()
31
32    def set_mask_format(self, mask_format: AttentionMaskFormat):
33        self.mask_format = mask_format
34
35    def set_mask_indice(self, mask, mask_index):
36        if mask in self.mask_indice:
37            assert mask_index == self.mask_indice[mask]
38        self.mask_indice[mask] = mask_index
39
40    def get_first_mask(self):
41        assert len(self.mask_indice) > 0
42        return next(iter(self.mask_indice))
43
44    def process_mask(self, mask_2d: str) -> str | None:
45        if self.mask_format == AttentionMaskFormat.NoMask:
46            return None
47
48        if mask_2d in self.mask_indice:
49            return self.mask_indice[mask_2d]
50
51        # Add cast to convert int64 to int32
52        if self.model.find_graph_input(mask_2d):
53            casted, input_name = self.utils.cast_graph_input_to_int32(mask_2d)
54        else:
55            input_name, _cast_node = self.utils.cast_input_to_int32(mask_2d)
56            casted = True
57
58        if casted:
59            self.mask_casted[mask_2d] = input_name
60
61        # Attention supports int32 attention mask (2D) since 1.4.0
62        if self.mask_format == AttentionMaskFormat.AttentionMask:
63            self.mask_indice[mask_2d] = input_name
64            return input_name
65
66        # Add a mask processing node to convert attention mask to mask index (1D)
67        output_name = self.model.create_node_name("mask_index")
68        if self.opset_version < 13:
69            mask_index_node = helper.make_node(
70                "ReduceSum",
71                inputs=[input_name],
72                outputs=[output_name],
73                name=self.model.create_node_name("ReduceSum", "MaskReduceSum"),
74            )
75            mask_index_node.attribute.extend([helper.make_attribute("axes", [1]), helper.make_attribute("keepdims", 0)])
76        else:
77            # ReduceSum-13: axes is moved from attribute to input
78            axes_name = "ort_const_1_reduce_sum_axes"
79            if self.model.get_initializer(axes_name) is None:
80                self.model.add_initializer(
81                    helper.make_tensor(
82                        name=axes_name,
83                        data_type=TensorProto.INT64,
84                        dims=[1],
85                        vals=[1],
86                        raw=False,
87                    )
88                )
89            mask_index_node = helper.make_node(
90                "ReduceSum",
91                inputs=[input_name, axes_name],
92                outputs=[output_name],
93                name=self.model.create_node_name("ReduceSum", "MaskReduceSum"),
94            )
95            mask_index_node.attribute.extend([helper.make_attribute("keepdims", 0)])
96
97        self.model.add_node(mask_index_node)
98
99        self.mask_indice[mask_2d] = output_name
100        return output_name
101
102
103class FusionAttention(Fusion):
104    """
105    Fuse Attention subgraph into one Attention node.
106    """
107
108    def __init__(
109        self,
110        model: OnnxModel,
111        hidden_size: int,
112        num_heads: int,
113        attention_mask: AttentionMask | None = None,
114        use_multi_head_attention: bool = False,
115        disable_multi_head_attention_bias: bool = False,
116        search_op_types: list[str] = ["SkipLayerNormalization", "LayerNormalization"],  # noqa: B006
117    ):
118        attention_op_name = "MultiHeadAttention" if use_multi_head_attention else "Attention"
119        super().__init__(model, attention_op_name, search_op_types)
120        self.hidden_size = hidden_size
121        self.num_heads = num_heads
122        self.attention_mask = attention_mask if attention_mask else AttentionMask(model)
123        self.use_multi_head_attention = use_multi_head_attention
124        self.disable_multi_head_attention_bias = disable_multi_head_attention_bias
125        self.mask_filter_value = None
126
127        # Flags to show warning only once
128        self.num_heads_warning = True
129        self.hidden_size_warning = True
130
131        self.shape_infer = None
132        self.shape_infer_done = True
133
134    def get_num_heads_and_hidden_size_from_concat(self, concat: NodeProto) -> tuple[int, int]:
135        """
136        Detect num_heads and hidden_size from Concat node in the following subgraph:
137
138        SkipLayerNormalization or EmbedLayerNormalization
139                        /        |
140                     MatMul    Shape
141                        |        |
142                       Add     Gather(indices=0)
143                        |        |
144                        |      Unsqueeze
145                        |        |
146                        |     Concat (*, -1, 12, 64)
147                        |     /
148                       Reshape
149                          |
150                       Transpose
151        """
152        if len(concat.input) == 4:
153            num_heads = self.model.get_constant_value(concat.input[2])
154            head_size = self.model.get_constant_value(concat.input[3])
155            if (
156                isinstance(num_heads, np.ndarray)
157                and num_heads.size == 1
158                and isinstance(head_size, np.ndarray)
159                and head_size.size == 1
160            ):
161                return num_heads[0], num_heads[0] * head_size[0]
162
163        return self.num_heads, self.hidden_size
164
165    def get_num_heads_and_hidden_size(self, reshape_q: NodeProto) -> tuple[int, int]:
166        """Detect num_heads and hidden_size from a reshape node.
167
168        Args:
169            reshape_q (NodeProto): reshape node for Q
170
171        Returns:
172            Tuple[int, int]: num_heads and hidden_size
173        """
174        # we assume that reshape fusion has done, so the shape is a tensor like [0, 0, num_heads, head_size]
175        q_shape_value = self.model.get_constant_value(reshape_q.input[1])
176        if q_shape_value is None:
177            concat = self.model.get_parent(reshape_q, 1)
178            if concat is not None and concat.op_type == "Concat":
179                return self.get_num_heads_and_hidden_size_from_concat(concat)
180            logger.debug("%s is not initializer.", reshape_q.input[1])
181            return self.num_heads, self.hidden_size  # Fall back to user specified value
182
183        if (
184            (not isinstance(q_shape_value, np.ndarray))
185            or len(q_shape_value) != 4
186            or (q_shape_value[2] <= 0 or q_shape_value[3] <= 0)
187        ):
188            logger.debug("q_shape_value=%s. Expected value are like [0, 0, num_heads, head_size].", q_shape_value)
189            return self.num_heads, self.hidden_size  # Fall back to user specified value
190
191        num_heads = q_shape_value[2]
192        head_size = q_shape_value[3]
193        hidden_size = num_heads * head_size
194
195        if self.num_heads > 0 and num_heads != self.num_heads:
196            if self.num_heads_warning:
197                logger.warning(
198                    "--num_heads is %d. Detected value is %d. Using detected value.", self.num_heads, num_heads
199                )
200                self.num_heads_warning = False  # Do not show the warning more than once
201
202        if self.hidden_size > 0 and hidden_size != self.hidden_size:
203            if self.hidden_size_warning:
204                logger.warning(
205                    "--hidden_size is %d. Detected value is %d. Using detected value.", self.hidden_size, hidden_size
206                )
207                self.hidden_size_warning = False  # Do not show the warning more than once
208
209        return num_heads, hidden_size
210
211    def get_add_qk_str(self, add_qk: NodeProto):
212        if not self.shape_infer_done:
213            self.shape_infer = self.model.infer_runtime_shape(update=True)
214            self.shape_infer_done = True
215
216        if self.shape_infer is None:
217            return None
218
219        input_0_shape = self.shape_infer.get_edge_shape(add_qk.input[0])
220        input_1_shape = self.shape_infer.get_edge_shape(add_qk.input[1])
221
222        if input_0_shape is None or input_1_shape is None:
223            logger.debug("one of the inputs of %s is None", add_qk)
224            return None
225
226        if input_0_shape != input_1_shape:
227            logger.debug("the shape of two inputs of %s is not same", add_qk)
228            return None
229
230        return add_qk.input[1]
231
232    def reshape_add_qk(self, add_qk: str):
233        # Convert 4D mask from (B,1,S,T) to (B,N,S,T)
234        # B = batch size, N = num heads, S = source sequence length, T = target sequence length
235        mask_output_name = add_qk + "_mask"
236
237        # Check if concat node for (B,1,S,T) --> (B,N,S,T) already exists
238        concat_node = list(filter(lambda node: node.output[0] == mask_output_name, self.nodes_to_add))
239        if len(concat_node) == 1:
240            return mask_output_name
241
242        assert len(concat_node) == 0
243        concat_node_name = self.model.create_node_name("Concat")
244        concat_add_qk_fp32 = helper.make_node(
245            "Concat",
246            inputs=[add_qk for _ in range(self.num_heads)],
247            outputs=[mask_output_name],
248            name=concat_node_name,
249            axis=1,
250        )
251        # Add new node to graph
252        self.nodes_to_add.append(concat_add_qk_fp32)
253        self.node_name_to_graph_name[concat_node_name] = self.this_graph_name
254
255        return mask_output_name
256
257    def concat_kv(self, past_k: str, past_v: str) -> str:
258        """Concatenate past_k and past_v inputs to create past_kv input.
259
260        Args:
261            past_k (str): name of past K value
262            past_v (str): name of past V value
263
264        Returns:
265            kv_output_name (str): name of past KV value
266        """
267        # Unsqueeze K and V nodes from (B,N,P,H) to (1,B,N,P,H)
268        # B = batch size, N = num heads, P = past sequence length, H = head size
269        unsqueeze_k_name = self.model.create_node_name("Unsqueeze")
270        unsqueeze_v_name = self.model.create_node_name("Unsqueeze")
271        k_5d_name = (past_k + "_5d").replace(".", "_")
272        v_5d_name = (past_v + "_5d").replace(".", "_")
273
274        k_5d = helper.make_node(
275            "Unsqueeze",
276            inputs=[past_k],
277            outputs=[k_5d_name],
278            name=unsqueeze_k_name,
279            axes=[0],
280        )
281        v_5d = helper.make_node(
282            "Unsqueeze",
283            inputs=[past_v],
284            outputs=[v_5d_name],
285            name=unsqueeze_v_name,
286            axes=[0],
287        )
288
289        # Add unsqueeze nodes to graph
290        self.nodes_to_add.append(k_5d)
291        self.nodes_to_add.append(v_5d)
292        self.node_name_to_graph_name[unsqueeze_k_name] = self.this_graph_name
293        self.node_name_to_graph_name[unsqueeze_v_name] = self.this_graph_name
294
295        # Concat K and V to get one node of size (2,B,N,P,H)
296        concat_node_name = self.model.create_node_name("Concat")
297        kv_output_name = past_v.replace(".value", ".kv").replace(".", "_").replace("_value", "_kv")
298        concat_kv = helper.make_node(
299            "Concat",
300            inputs=[k_5d_name, v_5d_name],
301            outputs=[kv_output_name],
302            name=concat_node_name,
303            axis=0,
304        )
305
306        # Add concat node to graph
307        self.nodes_to_add.append(concat_kv)
308        self.node_name_to_graph_name[concat_node_name] = self.this_graph_name
309
310        return kv_output_name
311
312    def split_kv(self, present_k_name: str, present_v_name: str, kv_node: str):
313        """Split kv_node containing present KV values into separate present K and present V values.
314
315        Args:
316            present_k_name (str): name of output to store present K value in
317            present_v_name (str): name of output to store present V value in
318            kv_node (str): name of present KV values
319        """
320        # Split kv_node into present_k and present_v nodes
321
322        # Create initializers for indexing kv_node, whose shape is (2,B,N,P,H)
323        k_index, v_index = "index_0", "index_1"
324        k_dim = self.model.get_initializer(k_index)
325        v_dim = self.model.get_initializer(v_index)
326        if k_dim is None:
327            k_dim = numpy_helper.from_array(np.array(0, dtype="int64"), name=k_index)
328            self.model.add_initializer(k_dim, self.this_graph_name)
329        if v_dim is None:
330            v_dim = numpy_helper.from_array(np.array(1, dtype="int64"), name=v_index)
331            self.model.add_initializer(v_dim, self.this_graph_name)
332
333        # Create nodes to index kv_node
334        gather_k_name = self.model.create_node_name("Gather")
335        gather_v_name = self.model.create_node_name("Gather")
336        present_k = helper.make_node(
337            "Gather",
338            inputs=[kv_node, k_index],
339            outputs=[present_k_name],
340            name=gather_k_name,
341            axis=0,
342        )
343        present_v = helper.make_node(
344            "Gather",
345            inputs=[kv_node, v_index],
346            outputs=[present_v_name],
347            name=gather_v_name,
348            axis=0,
349        )
350
351        # Add gather nodes to graph
352        self.nodes_to_add.append(present_k)
353        self.nodes_to_add.append(present_v)
354        self.node_name_to_graph_name[gather_k_name] = self.this_graph_name
355        self.node_name_to_graph_name[gather_v_name] = self.this_graph_name
356
357    def create_combined_qkv_bias(
358        self,
359        q_add: NodeProto,
360        k_add: NodeProto | None,
361        v_add: NodeProto | None,
362        name_prefix: str,
363    ) -> NodeProto | None:
364        q_bias = self.model.get_initializer(q_add.input[1]) or self.model.get_initializer(q_add.input[0])
365        qb = NumpyHelper.to_array(q_bias)
366        kb = np.zeros_like(qb)
367        vb = np.zeros_like(qb)
368        if k_add is not None:
369            k_bias = self.model.get_initializer(k_add.input[1]) or self.model.get_initializer(k_add.input[0])
370            kb = NumpyHelper.to_array(k_bias)
371        if v_add is not None:
372            v_bias = self.model.get_initializer(v_add.input[1]) or self.model.get_initializer(v_add.input[0])
373            vb = NumpyHelper.to_array(v_bias)
374
375        qkv_bias = np.stack((qb, kb, vb), axis=0)
376        qkv_bias_dim = 3 * np.prod(qb.shape)
377
378        bias_name = name_prefix + "_qkv_bias"
379        self.add_initializer(
380            name=bias_name,
381            data_type=q_bias.data_type,
382            dims=[qkv_bias_dim],
383            vals=qkv_bias,
384        )
385        return bias_name
386
387    def create_packed_qkv_matmul_node(
388        self,
389        q_matmul: NodeProto,
390        k_matmul: NodeProto,
391        v_matmul: NodeProto,
392        q_add: NodeProto,
393        k_add: NodeProto | None,
394        v_add: NodeProto | None,
395    ) -> tuple[NodeProto, NodeProto, NodeProto]:
396        """Create packed QKV MatMul node before MultiHeadAttention node.
397           This is for the scenario where an Attention node should be created but cannot be created
398           because past_key and past_value are separate inputs and not one concatenated input.
399
400        Args:
401            q_matmul (NodeProto): name of MatMul from Q path - (batch_size, sequence_length, hidden_size)
402            k_matmul (NodeProto): name of MatMul from K path - (batch_size, sequence_length, hidden_size)
403            v_matmul (NodeProto): name of MatMul from V path - (batch_size, sequence_length, hidden_size)
404            q_add (NodeProto): name of Add from Q path
405            k_add (NodeProto): name of Add from K path
406            v_add (NodeProto): name of Add from V path
407
408        Returns:
409             q_output (NodeProto): Slice node for Q
410             k_output (NodeProto): Slice node for K
411             v_output (NodeProto): Slice node for V
412        """
413        matmul_node_name = self.model.create_node_name("MatMul")
414
415        # Check that input for Q, K, V is the same
416        assert q_matmul.input[0] == k_matmul.input[0] and k_matmul.input[0] == v_matmul.input[0]
417
418        # Created packed QKV weight
419        q_weight = self.model.get_initializer(q_matmul.input[1])
420        k_weight = self.model.get_initializer(k_matmul.input[1])
421        v_weight = self.model.get_initializer(v_matmul.input[1])
422
423        qw = NumpyHelper.to_array(q_weight)
424        kw = NumpyHelper.to_array(k_weight)
425        vw = NumpyHelper.to_array(v_weight)
426
427        assert qw.shape == kw.shape and kw.shape == vw.shape
428        d = qw.shape[0]
429
430        qkv_weight = np.stack((qw, kw, vw), axis=1).reshape((d, 3 * d))
431        qkv_weight_name = matmul_node_name + "_qkv_weight"
432
433        self.add_initializer(
434            name=qkv_weight_name,
435            data_type=q_weight.data_type,
436            dims=[qkv_weight.shape[0], qkv_weight.shape[1]],
437            vals=qkv_weight,
438        )
439
440        # Created packed QKV MatMul with output (B, S, 3*D)
441        # Output is of the form:
442        #
443        # [[[Q Q ... Q Q K K ... K K V V ... V V]]]
444        #   [Q Q ... Q Q K K ... K K V V ... V V]
445        #                     .
446        #                     .
447        #                     .
448        #  [[Q Q ... Q Q K K ... K K V V ... V V]
449        #   [Q Q ... Q Q K K ... K K V V ... V V]]]
450        qkv_matmul_output = matmul_node_name + "_qkv_out"
451        qkv_matmul = helper.make_node(
452            "MatMul",
453            inputs=[q_matmul.input[0], qkv_weight_name],
454            outputs=[qkv_matmul_output],
455            name=matmul_node_name,
456        )
457        self.node_name_to_graph_name[matmul_node_name] = self.this_graph_name
458
459        qkv_nodes = [qkv_matmul]
460
461        # Create Slice nodes to access Q, K, V
462        q_slice_name = matmul_node_name + "_q_start_index"
463        self.add_initializer(name=q_slice_name, data_type=TensorProto.INT64, dims=[1], vals=[0], raw=False)
464        k_slice_name = matmul_node_name + "_k_start_index"
465        self.add_initializer(name=k_slice_name, data_type=TensorProto.INT64, dims=[1], vals=[d], raw=False)
466        v_slice_name = matmul_node_name + "_v_start_index"
467        self.add_initializer(name=v_slice_name, data_type=TensorProto.INT64, dims=[1], vals=[2 * d], raw=False)
468        end_of_qkv_name = matmul_node_name + "_end_of_qkv_index"
469        self.add_initializer(name=end_of_qkv_name, data_type=TensorProto.INT64, dims=[1], vals=[3 * d], raw=False)
470        qkv_last_axis_name = matmul_node_name + "_qkv_last_axis"
471        self.add_initializer(name=qkv_last_axis_name, data_type=TensorProto.INT64, dims=[1], vals=[-1], raw=False)
472
473        q_slice_output = matmul_node_name + "_q_out"
474        q_slice = helper.make_node(
475            "Slice",
476            inputs=[qkv_matmul_output, q_slice_name, k_slice_name, qkv_last_axis_name],
477            outputs=[q_slice_output],
478            name=self.model.create_node_name("Slice"),
479        )
480        self.node_name_to_graph_name[q_slice.name] = self.this_graph_name
481        k_slice_output = matmul_node_name + "_k_out"
482        k_slice = helper.make_node(
483            "Slice",
484            inputs=[qkv_matmul_output, k_slice_name, v_slice_name, qkv_last_axis_name],
485            outputs=[k_slice_output],
486            name=self.model.create_node_name("Slice"),
487        )
488        self.node_name_to_graph_name[k_slice.name] = self.this_graph_name
489        v_slice_output = matmul_node_name + "_v_out"
490        v_slice = helper.make_node(
491            "Slice",
492            inputs=[qkv_matmul_output, v_slice_name, end_of_qkv_name, qkv_last_axis_name],
493            outputs=[v_slice_output],
494            name=self.model.create_node_name("Slice"),
495        )
496        self.node_name_to_graph_name[v_slice.name] = self.this_graph_name
497
498        q_output = q_slice
499        k_output = k_slice
500        v_output = v_slice
501        qkv_nodes.extend([q_slice, k_slice, v_slice])
502
503        if self.disable_multi_head_attention_bias:
504            if q_add is not None:
505                initializer_input = 1 if self.model.get_initializer(q_add.input[1]) else 0
506                if np.any(NumpyHelper.to_array(self.model.get_initializer(q_add.input[initializer_input]))):
507                    q_add.input[1 - initializer_input] = q_slice_output
508                    q_output = q_add
509                    qkv_nodes.append(q_add)
510                    self.node_name_to_graph_name[q_add.name] = self.this_graph_name
511            if k_add is not None:
512                initializer_input = 1 if self.model.get_initializer(k_add.input[1]) else 0
513                if np.any(NumpyHelper.to_array(self.model.get_initializer(k_add.input[initializer_input]))):
514                    k_add.input[1 - initializer_input] = k_slice_output
515                    k_output = k_add
516                    qkv_nodes.append(k_add)
517                    self.node_name_to_graph_name[k_add.name] = self.this_graph_name
518            if v_add is not None:
519                initializer_input = 1 if self.model.get_initializer(v_add.input[1]) else 0
520                if np.any(NumpyHelper.to_array(self.model.get_initializer(v_add.input[initializer_input]))):
521                    v_add.input[1 - initializer_input] = v_slice_output
522                    v_output = v_add
523                    qkv_nodes.append(v_add)
524                    self.node_name_to_graph_name[v_add.name] = self.this_graph_name
525
526        # Add nodes to graph
527        self.nodes_to_add.extend(qkv_nodes)
528        return q_output, k_output, v_output
529
530    # This function is used in child classes for bart or conformer model.
531    def create_multihead_attention_node(
532        self,
533        q_matmul: NodeProto,
534        k_matmul: NodeProto | str | None,
535        v_matmul: NodeProto | str | None,
536        q_add: NodeProto,
537        k_add: NodeProto | None,
538        v_add: NodeProto | None,
539        num_heads: int,
540        hidden_size: int,
541        output: str,
542        key_padding_mask: str = "",
543        add_qk: str = "",
544        unidirectional: bool = False,
545        past_k: str = "",
546        past_v: str = "",
547        present_k: str = "",
548        present_v: str = "",
549        packed_qkv: bool = False,
550    ) -> NodeProto | None:
551        """Create a MultiHeadAttention node.
552
553        Args:
554            q_matmul (NodeProto): name of MatMul from Q path - (batch_size, sequence_length, hidden_size)
555            k_matmul (NodeProto): name of MatMul from K path - (batch_size, sequence_length, hidden_size) or (batch_size, num_heads, past_sequence_length, head_size)
556            v_matmul (NodeProto): name of MatMul from V path - (batch_size, sequence_length, hidden_size) or (batch_size, num_heads, past_sequence_length, head_size)
557            q_add (NodeProto): name of Add from Q path
558            k_add (NodeProto): name of Add from K path
559            v_add (NodeProto): name of Add from V path
560            num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
561            hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
562            output (str): output name of MHA
563            key_padding_mask (str): name of key padding mask
564            add_qk (str): name of add after Q x K'
565            unidirectional (bool): whether to apply causal attention mask automatically or not
566            past_k (str): name of past K value - (batch_size, num_heads, past_sequence_length, head_size)
567            past_v (str): name of past V value - (batch_size, num_heads, past_sequence_length, head_size)
568            present_k (str): name of present K value - (batch_size, num_heads, sequence_length, head_size)
569            present_v (str): name of present V value - (batch_size, num_heads, sequence_length, head_size)
570            packed_qkv (bool): whether to combine MatMuls from Q, K, V paths
571                               Note: This is for the scenario where an Attention node should be created but cannot be created
572                               because past_key and past_value are separate inputs and not one concatenated input.
573
574        Returns:
575            Union[NodeProto, None]: the node created or None if failed.
576        """
577        # B = batch size, N = num heads, P = past seq len, H = head size
578        assert num_heads > 0
579
580        if hidden_size > 0 and (hidden_size % num_heads) != 0:
581            logger.debug("input hidden size %d is not a multiple of num of heads %d", hidden_size, num_heads)
582            return None
583
584        graph_input_names = {node.name for node in self.model.graph().input}
585        mha_node_name = self.model.create_node_name("Attention")
586
587        # Add initial Q/K/V inputs for MHA
588        mha_inputs = []
589        if packed_qkv:
590            q_slice, k_slice, v_slice = self.create_packed_qkv_matmul_node(
591                q_matmul,
592                k_matmul,
593                v_matmul,
594                q_add,
595                k_add,
596                v_add,
597            )
598            mha_inputs.extend([q_slice.output[0], k_slice.output[0], v_slice.output[0]])
599        elif isinstance(k_matmul, NodeProto) and isinstance(v_matmul, NodeProto):
600            if self.disable_multi_head_attention_bias:
601                mha_inputs.extend([q_add.output[0], k_matmul.output[0], v_add.output[0]])
602            else:
603                mha_inputs.extend([q_matmul.output[0], k_matmul.output[0], v_matmul.output[0]])
604        elif (
605            isinstance(k_matmul, str)
606            and isinstance(v_matmul, str)
607            and k_matmul in graph_input_names
608            and v_matmul in graph_input_names
609        ):
610            if self.disable_multi_head_attention_bias:
611                mha_inputs.extend([q_add.output[0], k_matmul, v_matmul])
612            else:
613                mha_inputs.extend([q_matmul.output[0], k_matmul, v_matmul])
614        else:
615            return None
616
617        # Add bias to inputs for MHA
618        # Bias for cross attention is not fully supported in DMMHA and cpu MHA kernels since they assume
619        # bias has been added to key and value when they are in BNSH format, so only bias for query is used.
620        # Need add checks if we found such assumption is not true.
621        if not self.disable_multi_head_attention_bias:
622            bias_name = self.create_combined_qkv_bias(q_add, k_add, v_add, mha_node_name)
623            mha_inputs.append(bias_name)
624        else:
625            mha_inputs.append("")
626
627        # Add optional inputs for MHA
628        if past_k and past_v:
629            mha_inputs.extend([key_padding_mask, add_qk, past_k, past_v])
630        elif key_padding_mask or add_qk:
631            mha_inputs.extend([key_padding_mask, add_qk])
632
633        # Add outputs for MHA
634        mha_outputs = [output]
635        if present_k and present_v:
636            mha_outputs.extend([present_k, present_v])
637
638        mha_node = helper.make_node(
639            "MultiHeadAttention",
640            inputs=mha_inputs,
641            outputs=mha_outputs,
642            name=mha_node_name,
643        )
644        mha_node.domain = "com.microsoft"
645        mha_node.attribute.append(helper.make_attribute("num_heads", num_heads))
646        if unidirectional:
647            mha_node.attribute.append(helper.make_attribute("unidirectional", int(unidirectional)))
648
649        self.increase_counter("MultiHeadAttention")
650        return mha_node
651
652    def create_attention_node(
653        self,
654        mask_index: str | None,
655        q_matmul: NodeProto,
656        k_matmul: NodeProto,
657        v_matmul: NodeProto,
658        q_add: NodeProto,
659        k_add: NodeProto,
660        v_add: NodeProto,
661        num_heads: int,
662        hidden_size: int,
663        first_input: str,
664        output: str,
665        add_qk_str: str = "",
666        causal: bool = False,
667        past_k: str = "",
668        past_v: str = "",
669        present_k: str = "",
670        present_v: str = "",
671        scale: float | None = None,
672    ) -> NodeProto | None:
673        """Create an Attention node.
674
675        Args:
676            mask_index (str | None): mask input
677            q_matmul (NodeProto): MatMul node in fully connection for Q
678            k_matmul (NodeProto): MatMul node in fully connection for K
679            v_matmul (NodeProto): MatMul node in fully connection for V
680            q_add (NodeProto): Add bias node in fully connection for Q
681            k_add (NodeProto): Add bias node in fully connection for K
682            v_add (NodeProto): Add bias node in fully connection for V
683            num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
684            hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
685            first_input (str): first input name
686            output (str): output name
687            add_qk_str (str): name of Add node after Q x K'
688            causal: whether it is uni-directional mask.
689            past_k (str): name of input for past K value
690            past_v (str): name of input for past V value
691            present_k (str): name of output to store present K value
692            present_v (str): name of output to store present V value
693            scale: scale before softmax
694
695        Returns:
696            Union[NodeProto, None]: the node created or None if failed.
697        """
698        assert num_heads > 0
699
700        if hidden_size > 0 and (hidden_size % num_heads) != 0:
701            logger.debug("input hidden size %d is not a multiple of num of heads %d", hidden_size, num_heads)
702            return None
703
704        has_bias = True
705        if q_add is None and k_add is None and v_add is None:
706            has_bias = False
707
708        q_weight = self.model.get_initializer(q_matmul.input[1])
709        k_weight = self.model.get_initializer(k_matmul.input[1])
710        v_weight = self.model.get_initializer(v_matmul.input[1])
711
712        q_bias, k_bias, v_bias = None, None, None
713        if has_bias:
714            q_bias = self.model.get_initializer(q_add.input[1]) or self.model.get_initializer(q_add.input[0])
715            k_bias = self.model.get_initializer(k_add.input[1]) or self.model.get_initializer(k_add.input[0])
716            v_bias = self.model.get_initializer(v_add.input[1]) or self.model.get_initializer(v_add.input[0])
717
718            if not (k_weight and v_weight and q_bias and k_bias):
719                return None
720
721        if q_weight is None:
722            print(
723                f"{q_matmul.input[1]} is not an initializer. "
724                "Please set do_constant_folding=True in torch.onnx.export to unblock attention fusion"
725            )
726            return None
727
728        qw = NumpyHelper.to_array(q_weight)
729        kw = NumpyHelper.to_array(k_weight)
730        vw = NumpyHelper.to_array(v_weight)
731
732        # assert q and k have same shape as expected
733        assert qw.shape == kw.shape
734
735        qw_in_size = qw.shape[0]
736        kw_in_size = kw.shape[0]
737        vw_in_size = vw.shape[0]
738
739        assert qw_in_size == kw_in_size == vw_in_size
740
741        if hidden_size > 0 and hidden_size != qw_in_size:
742            logger.warning(
743                "Input hidden size (%d) is not same as weight matrix dimension of q,k,v (%d). "
744                "Please provide a correct input hidden size or pass in 0",
745                hidden_size,
746                qw_in_size,
747            )
748
749        is_qkv_diff_dims = False
750        if qw.shape != vw.shape:
751            is_qkv_diff_dims = True
752
753        # All the matrices can have the same shape or q, k matrices can have the same shape with v being different
754        # For 2d weights, the shapes would be [in_size, out_size].
755        # For 3d weights, shape would be [in_size, a, b] where a*b = out_size
756        qw_out_size = np.prod(qw.shape[1:])
757        kw_out_size = np.prod(kw.shape[1:])
758        vw_out_size = np.prod(vw.shape[1:])
759
760        qkv_weight_dim = 0
761        if is_qkv_diff_dims:
762            qkv_weight = np.concatenate((qw, kw, vw), axis=1)
763            qkv_weight_dim = qw_out_size + kw_out_size + vw_out_size
764        else:
765            qkv_weight = np.stack((qw, kw, vw), axis=1)
766            qkv_weight_dim = 3 * qw_out_size
767
768        qkv_bias_dim = 0
769        qkv_bias: np.ndarray | None = None
770        if has_bias:
771            qb = NumpyHelper.to_array(q_bias)
772            kb = NumpyHelper.to_array(k_bias)
773            vb = NumpyHelper.to_array(v_bias)
774
775            q_bias_shape = np.prod(qb.shape)
776            k_bias_shape = np.prod(kb.shape)
777            v_bias_shape = np.prod(vb.shape)
778
779            assert q_bias_shape == k_bias_shape == qw_out_size
780            assert v_bias_shape == vw_out_size
781
782            if is_qkv_diff_dims:
783                qkv_bias = np.concatenate((qb, kb, vb), axis=0)
784                qkv_bias_dim = q_bias_shape + k_bias_shape + v_bias_shape
785            else:
786                qkv_bias = np.stack((qb, kb, vb), axis=0)
787                qkv_bias_dim = 3 * q_bias_shape
788
789        attention_node_name = self.model.create_node_name("Attention")
790
791        if not self.use_multi_head_attention:
792            self.add_initializer(
793                name=attention_node_name + "_qkv_weight",
794                data_type=q_weight.data_type,
795                dims=[qw_in_size, int(qkv_weight_dim)],
796                vals=qkv_weight,
797            )
798
799        if has_bias:
800            self.add_initializer(
801                name=attention_node_name + "_qkv_bias",
802                data_type=q_bias.data_type,
803                dims=[int(qkv_bias_dim)],
804                vals=qkv_bias,
805            )
806
807        # For MultiHeadAttention operator, use separated inputs for query, key and value, and no weights.
808        if self.use_multi_head_attention:
809            if add_qk_str:
810                logger.debug("MultiHeadAttention does not support relative_position_bias: cannot fuse the attention.")
811                return None
812
813            attention_inputs = [
814                q_matmul.output[0],
815                k_matmul.output[0],
816                v_matmul.output[0],
817                attention_node_name + "_qkv_bias",
818            ]
819
820            if mask_index is not None:
821                attention_inputs.append(mask_index)
822
823            attention_node = helper.make_node(
824                "MultiHeadAttention",
825                inputs=attention_inputs,
826                outputs=[output],
827                name=attention_node_name,
828            )
829            self.increase_counter("MultiHeadAttention")
830
831        else:
832            attention_inputs = [
833                first_input,
834                attention_node_name + "_qkv_weight",
835                attention_node_name + "_qkv_bias" if has_bias else "",
836            ]
837            if mask_index is not None:
838                attention_inputs.append(mask_index)
839            else:
840                attention_inputs.append("")
841
842            past_exists = past_k and past_v
843            if past_exists:
844                past_kv = self.concat_kv(past_k, past_v)
845                attention_inputs.append(past_kv)
846
847            if add_qk_str:
848                # Add additional add to attention node (input name = attention_bias)
849                if not past_exists:
850                    attention_inputs.append("")
851                attention_inputs.append(add_qk_str)
852
853            attention_outputs = [output]
854            if present_k and present_v:
855                present_kv = present_k.replace(".key", "").replace("_key", "").replace(".", "_")
856                attention_outputs.append(present_kv)
857                self.split_kv(present_k, present_v, present_kv)
858
859            attention_node = helper.make_node(
860                "Attention",
861                inputs=attention_inputs,
862                outputs=attention_outputs,
863                name=attention_node_name,
864            )
865            self.increase_counter("Attention")
866
867        attention_node.domain = "com.microsoft"
868        attention_node.attribute.extend([helper.make_attribute("num_heads", num_heads)])
869
870        if causal:
871            attention_node.attribute.extend([helper.make_attribute("unidirectional", 1)])
872
873        if scale is not None:
874            attention_node.attribute.extend([helper.make_attribute("scale", scale)])
875
876        if is_qkv_diff_dims:
877            attention_node.attribute.extend(
878                [helper.make_attribute("qkv_hidden_sizes", [qw_out_size, kw_out_size, vw_out_size])]
879            )
880
881        if self.mask_filter_value is not None:
882            attention_node.attribute.extend([helper.make_attribute("mask_filter_value", float(self.mask_filter_value))])
883
884        return attention_node
885
886    def fuse(self, node, input_name_to_nodes, output_name_to_node):
887        # Sometimes we can not fuse skiplayernormalization since the add before layernorm has an output that used by nodes outside skiplayernorm
888        # Conceptually we treat add before layernorm as skiplayernorm node since they share the same pattern
889        normalize_node = node
890        start_node = normalize_node
891        if normalize_node.op_type == "LayerNormalization":
892            add_before_layernorm = self.model.match_parent(normalize_node, "Add", 0)
893            if add_before_layernorm is not None:
894                start_node = add_before_layernorm
895            elif self.model.find_graph_input(normalize_node.input[0]) is not None:
896                # Pre-LN first block: LN fed directly by graph input.  QKV matching will
897                # still fail from this (first) LN anchor because its inputs are weights, not
898                # the QKV projection path.  The real fusion happens when fuse() is called
899                # again from the second LN/SkipLN anchor after the residual Add, where the
900                # other_inputs and root_input changes (#2-#4) take effect.
901                start_node = normalize_node
902            else:
903                return
904
905        # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
906        qkv_nodes = self.model.match_parent_path(
907            start_node,
908            ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
909            [None, None, 0, 0, 0],
910        )
911        einsum_node = None
912        if qkv_nodes is not None:
913            (_, _, reshape_qkv, transpose_qkv, matmul_qkv) = qkv_nodes
914        else:
915            # Match Albert
916            qkv_nodes = self.model.match_parent_path(
917                start_node, ["Add", "Einsum", "Transpose", "MatMul"], [1, None, 0, 0]
918            )
919            if qkv_nodes is not None:
920                (_, einsum_node, transpose_qkv, matmul_qkv) = qkv_nodes
921            else:
922                return
923
924        other_inputs = []
925        for _i, node_input in enumerate(start_node.input):
926            if node_input not in output_name_to_node:
927                if self.model.find_graph_input(node_input) is None:
928                    continue
929
930            if node_input == qkv_nodes[0].output[0]:
931                continue
932            other_inputs.append(node_input)
933        if len(other_inputs) != 1:
934            return
935
936        root_input = other_inputs[0]
937
938        # Match flaubert                     Mask
939        #                                     |
940        # Mul --> LayerNormalization -->  Attention --> MatMul --> Add
941        #  |                                                        |
942        #  |                                                        |
943        #  +---------------------------------------------------------
944        mul_before_layernorm = self.model.match_parent(start_node, "Mul", 0)
945        if mul_before_layernorm is not None:
946            mul_children = input_name_to_nodes[mul_before_layernorm.output[0]]
947            if mul_children is not None and len(mul_children) == 2:
948                layernorm_node = mul_children[1]
949                if layernorm_node.op_type == "LayerNormalization":
950                    root_input = layernorm_node.output[0]
951                else:
952                    return
953            elif mul_children is not None and len(mul_children) == 5:
954                root_input = mul_before_layernorm.output[0]
955            else:
956                return
957        elif normalize_node.op_type in ("LayerNormalization", "SkipLayerNormalization"):
958            children = input_name_to_nodes[root_input]
959            for child in children:
960                if child.op_type == "LayerNormalization":
961                    root_input = child.output[0]
962
963        # When Add before the LayerNormalization produces an output
964        # that is consumed by some other nodes other than the LayerNormalization itself,
965        # fused SkipLayerNormalization will have several outputs.
966        # In this case we need to pick the one used in Attention
967        # For example, this is the case for ViT
968        # SkipLayerNormalization --> Attention --> MatMul --> Add --> SkipLayerNormalization
969        #  |                                                                     |
970        #  |                                                                     |
971        #  +---------------------------------------------------------------------+
972        if root_input in output_name_to_node:
973            parent_node = output_name_to_node[root_input]
974            if parent_node.op_type == "SkipLayerNormalization" and len(parent_node.output) == 4:
975                root_input = parent_node.output[0]
976
977        children = input_name_to_nodes[root_input]
978        children_types = [child.op_type for child in children]
979        if children_types.count("MatMul") != 3:
980            return
981
982        v_nodes = self.model.match_parent_path(matmul_qkv, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None])
983        if v_nodes is None:
984            logger.debug("fuse_attention: failed to match v path")
985            return
986        (_, _, add_v, matmul_v) = v_nodes
987
988        is_distill = False
989        is_distill_add = False
990        is_no_mask_attention = False
991        is_sdpa = False
992        qk_paths = {
993            "path1": (["Softmax", "Add", "Div", "MatMul"], [0, 0, None, 0]),
994            "path2": (["Softmax", "Add", "Mul", "MatMul"], [0, 0, None, 0]),
995            "path3": (["Softmax", "Where", "MatMul", "Div"], [0, 0, 2, 0]),
996            "path4": (["Softmax", "Add", "Where", "MatMul"], [0, 0, 0, 2]),
997            "path5": (["Softmax", "Div", "MatMul"], [0, 0, 0]),
998            "sdpa": (["Softmax", "Add", "MatMul", "Mul", "Sqrt"], [0, 0, None, 0, 1]),
999        }
1000
1001        qk_nodes = None
1002        for k, v in qk_paths.items():
1003            qk_nodes = self.model.match_parent_path(matmul_qkv, v[0], v[1])
1004            if qk_nodes is None:
1005                continue
1006            if k == "path3":
1007                is_distill = True
1008            elif k == "path4":
1009                is_distill_add = True
1010            elif k == "path5":
1011                is_no_mask_attention = True
1012            elif k == "sdpa":
1013                is_sdpa = True
1014            break
1015
1016        if qk_nodes is None:
1017            logger.debug("fuse_attention: failed to match qk path")
1018            return
1019
1020        add_qk = None
1021        matmul_qk = None
1022        where_qk = None
1023        after_q = None
1024        if is_distill:
1025            (_, where_qk, matmul_qk, _) = qk_nodes
1026        elif is_distill_add:
1027            (_, add_qk, where_qk, matmul_qk) = qk_nodes
1028        elif is_no_mask_attention:
1029            (_, _, matmul_qk) = qk_nodes
1030        elif is_sdpa:
1031            (_, add_qk, matmul_qk, after_q, _) = qk_nodes
1032        else:
1033            (_, add_qk, _, matmul_qk) = qk_nodes
1034
1035        after_q = after_q or matmul_qk
1036        q_nodes = self.model.match_parent_path(after_q, ["Transpose", "Reshape", "Add", "MatMul"], [0, 0, 0, None])
1037        if q_nodes is None:
1038            q_nodes = self.model.match_parent_path(
1039                after_q,
1040                ["Div", "Transpose", "Reshape", "Add", "MatMul"],
1041                [0, 0, 0, 0, None],
1042            )
1043            if q_nodes is None:
1044                logger.debug("fuse_attention: failed to match q path")
1045                return
1046        reshape_q = q_nodes[-3]
1047        add_q = q_nodes[-2]
1048        matmul_q = q_nodes[-1]
1049
1050        after_k = matmul_qk
1051        if is_sdpa:
1052            mul_k_nodes = self.model.match_parent_path(matmul_qk, ["Mul", "Sqrt"], [1, None])
1053            if mul_k_nodes is None:
1054                logger.debug("fuse_attention: failed to match mul sqrt q path")
1055                return
1056            (after_k, _) = mul_k_nodes
1057
1058        k_nodes = self.model.match_parent_path(
1059            after_k, ["Transpose", "Reshape", "Add", "MatMul"], [0 if is_sdpa else 1, 0, 0, None]
1060        )
1061        if k_nodes is None:
1062            k_nodes = self.model.match_parent_path(
1063                matmul_qk,
1064                ["Transpose", "Transpose", "Reshape", "Add", "MatMul"],
1065                [1, 0, 0, 0, None],
1066            )
1067            if k_nodes is None:
1068                logger.debug("fuse_attention: failed to match k path")
1069                return
1070        add_k = k_nodes[-2]
1071        matmul_k = k_nodes[-1]
1072
1073        # Note that Cast might be removed by OnnxRuntime so we match two patterns here.
1074        mask_nodes = None
1075        add_qk_str = ""
1076        if is_distill:
1077            _, mask_nodes, _ = self.model.match_parent_paths(
1078                where_qk,
1079                [
1080                    (["Expand", "Reshape", "Equal"], [0, 0, 0]),
1081                    (["Equal", "Unsqueeze", "Unsqueeze"], [0, 0, 0]),
1082                    (["Cast", "Expand", "Reshape", "Equal"], [0, 0, 0, 0]),
1083                ],
1084                output_name_to_node,
1085            )
1086        elif is_distill_add:
1087            _, mask_nodes, _ = self.model.match_parent_paths(
1088                where_qk,
1089                [
1090                    (["Cast", "Equal", "Unsqueeze", "Unsqueeze"], [0, 0, 0, 0]),
1091                    (["Equal", "Unsqueeze", "Unsqueeze"], [0, 0, 0]),
1092                ],
1093                output_name_to_node,
1094            )
1095            if add_qk is not None:
1096                add_qk_str = self.get_add_qk_str(add_qk)
1097                if add_qk_str is None:
1098                    logger.debug("fuse_attention: failed to verify shape inference of %s", add_qk)
1099                    return
1100        elif is_no_mask_attention:
1101            pass
1102        else:
1103            _, mask_nodes, _ = self.model.match_parent_paths(
1104                add_qk,
1105                [
1106                    (["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze"], [None, 0, 1, 0, 0]),
1107                    (["Mul", "Sub", "Unsqueeze", "Unsqueeze"], [None, 0, 1, 0]),
1108                    # The following two patterns are for SDPA.
1109                    (["Where", "Cast", "Sub", "Expand", "Unsqueeze", "Unsqueeze"], [None, 0, 0, 1, 0, 0]),
1110                    (["Where", "Cast", "Sub", "Cast", "Expand", "Unsqueeze", "Unsqueeze"], [None, 0, 0, 1, 0, 0, 0]),
1111                ],
1112                output_name_to_node,
1113            )
1114        if not is_no_mask_attention and mask_nodes is None:
1115            logger.debug("fuse_attention: failed to match mask path")
1116            return
1117
1118        if not is_no_mask_attention and len(mask_nodes) > 1:
1119            _, mul_val = self.model.get_constant_input(mask_nodes[0])
1120            # The mask value shall be a float scalar (usually is the lowest float value).
1121            if (
1122                (mul_val is None)
1123                or not (isinstance(mul_val, np.ndarray) and mul_val.size == 1)
1124                or (mul_val.item() >= 0)
1125            ):
1126                return
1127            if mul_val.item() != -10000:
1128                self.mask_filter_value = mul_val.item()
1129
1130        if matmul_v.input[0] == root_input and matmul_q.input[0] == root_input and matmul_k.input[0] == root_input:
1131            mask_index = self.attention_mask.process_mask(mask_nodes[-1].input[0]) if not is_no_mask_attention else None
1132
1133            attention_last_node = reshape_qkv if einsum_node is None else transpose_qkv
1134
1135            q_num_heads, q_hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
1136            if q_num_heads <= 0 or q_hidden_size <= 0:
1137                logger.warning(
1138                    "Failed to detect num_heads and hidden_size for Attention fusion. "
1139                    "Please specify those parameters in argument."
1140                )
1141                return
1142
1143            # number of heads are same for all the paths, hence to create attention node, we pass the q_num_heads
1144            # the input_hidden_size represents the input hidden size, this is used as needed but hidden sizes for Q, K are extracted appropriately
1145            new_node = self.create_attention_node(
1146                mask_index=mask_index,
1147                q_matmul=matmul_q,
1148                k_matmul=matmul_k,
1149                v_matmul=matmul_v,
1150                q_add=add_q,
1151                k_add=add_k,
1152                v_add=add_v,
1153                num_heads=q_num_heads,
1154                hidden_size=q_hidden_size,
1155                first_input=root_input,
1156                output=attention_last_node.output[0],
1157                add_qk_str=add_qk_str,
1158            )
1159
1160            if new_node is None:
1161                return
1162
1163            self.nodes_to_add.append(new_node)
1164            self.node_name_to_graph_name[new_node.name] = self.this_graph_name
1165
1166            if einsum_node is not None:
1167                unique_index = einsum_node.input[0]
1168                new_edge = "edge_modified_" + unique_index
1169
1170                shape_tensor = self.add_initializer(
1171                    name="shape_modified_tensor" + unique_index,
1172                    data_type=TensorProto.INT64,
1173                    dims=[4],
1174                    vals=[0, 0, q_num_heads, int(q_hidden_size / q_num_heads)],
1175                    raw=False,
1176                )
1177
1178                self.model.add_node(
1179                    helper.make_node(
1180                        "Reshape",
1181                        [attention_last_node.output[0], shape_tensor.name],
1182                        [new_edge],
1183                        "reshape_modified_" + unique_index,
1184                    ),
1185                    self.this_graph_name,
1186                )
1187                einsum_node.input[0] = new_edge
1188
1189            self.nodes_to_remove.extend([attention_last_node, transpose_qkv, matmul_qkv])
1190            self.nodes_to_remove.extend(qk_nodes)
1191
1192            # For MultiHeadAttention operator, MatMul nodes for Q/K/V projection shall not be fused.
1193            self.nodes_to_remove.extend(q_nodes if not self.use_multi_head_attention else q_nodes[:-1])
1194            self.nodes_to_remove.extend(k_nodes if not self.use_multi_head_attention else k_nodes[:-1])
1195            self.nodes_to_remove.extend(v_nodes if not self.use_multi_head_attention else v_nodes[:-1])
1196
1197            # Use prune graph to remove mask nodes since they are shared by all attention nodes.
1198            self.prune_graph = True
1199 
codekingpro/portable-devtools · Team Ai