Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_gpt_attention.py547 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 FusionUtils
10from onnx import helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionGptAttentionPastBase(Fusion):
17    """Base class for GPT Attention Fusion with past state"""
18
19    def __init__(self, model: OnnxModel, num_heads: int):
20        super().__init__(model, "Attention", ["LayerNormalization", "SkipLayerNormalization"], "with past")
21        self.num_heads = num_heads
22        self.utils = FusionUtils(model)
23        self.casted_attention_mask = {}  # map from name of attention mask to the name that casted to int32
24        self.mask_filter_value = None
25
26    def match_past_pattern_1(self, concat_k, concat_v, output_name_to_node):
27        # Pattern 1:
28        #                      {past}
29        #                    /        \
30        #                   /          \
31        #    Gather(axes=0, indices=0)  Gather(indices=1)
32        #      |                          |
33        #    Transpose (perm=0,1,3,2)     |
34        #      |                          |
35        #  Concat_k                     Concat_v
36        #      |                        /
37        #  Transpose (perm=0,1,3,2)    /
38        #      |                      /
39        #  Unsqueeze        Unsqueeze
40        #        \        /
41        #         \      /
42        #           Concat
43        #             |
44        #         {present}
45        gather = self.model.get_parent(concat_v, 0, output_name_to_node)
46        if gather is None or gather.op_type != "Gather":
47            logger.debug("match_past_pattern_1: expect Gather for past")
48            return None
49
50        if self.model.find_constant_input(gather, 1) != 1:
51            logger.debug("match_past_pattern_1: expect indices=1 for Gather of past")
52            return None
53        past = gather.input[0]
54
55        parent = self.model.get_parent(concat_k, 0, output_name_to_node)
56        if parent and parent.op_type == "Gather":
57            gather_past_k = parent
58        else:
59            past_k_nodes = self.model.match_parent_path(concat_k, ["Transpose", "Gather"], [0, 0])
60            if past_k_nodes is None:
61                logger.debug("match_past_pattern_1: failed match Transpose and Gather")
62                return None
63            gather_past_k = past_k_nodes[-1]
64
65        if self.model.find_constant_input(gather_past_k, 0) != 1:
66            logger.debug("match_past_pattern_1: expect indices=0 for Gather k of past")
67            return None
68        past_k = gather_past_k.input[0]
69        if past != past_k:
70            logger.debug("match_past_pattern_1: expect past to be same")
71            return None
72
73        return past
74
75    def match_past_pattern_2(self, concat_k, concat_v, output_name_to_node):
76        # Pattern 2:
77        #      Split (QKV)
78        #      / |   |
79        #     /  |   +----------------------+
80        #        |                          |
81        #        |         {past}           |
82        #        |           |              |
83        #      Reshape     Split         Reshape
84        #        |         /    \           |
85        # Transpose_k  Squeeze  Squeeze  Transpose_v
86        #        |      |        \        /
87        #        +------|---+     \      /
88        #               |   |      \    /
89        #              Concat_k   Concat_v
90        #               |            |
91        #          Unsqueeze    Unsqueeze
92        #                \       /
93        #                 Concat
94        #                   |
95        #               {present}
96        #
97        squeeze = self.model.get_parent(concat_v, 0, output_name_to_node)
98        if squeeze is None or squeeze.op_type != "Squeeze":
99            logger.debug("match_past_pattern_2: expect Squeeze as parent of concat_v")
100            return None
101
102        split = self.model.get_parent(squeeze, 0, output_name_to_node)
103        if split is None or split.op_type != "Split":
104            logger.debug("match_past_pattern_2: expect Split for past path")
105            return None
106
107        opset_version = self.model.get_opset_version()
108        if opset_version < 13:
109            if not FusionUtils.check_node_attribute(squeeze, "axes", [0]):
110                logger.debug("match_past_pattern_2: axes != [0] for Squeeze in past path")
111                return None
112
113            if not FusionUtils.check_node_attribute(split, "split", [1, 1]):
114                logger.debug("match_past_pattern_2: split != [1, 1] for Split in past path")
115                return None
116        else:
117            if not self.utils.check_node_input_value(squeeze, 1, [0]):
118                logger.debug("match_past_pattern_2: axes != [0] for Squeeze in past path")
119                return None
120
121            if not self.utils.check_node_input_value(split, 1, [1, 1]):
122                logger.debug("match_past_pattern_2: split != [1, 1] for Split in past path")
123                return None
124
125        if not FusionUtils.check_node_attribute(split, "axis", 0, default_value=0):
126            logger.debug("match_past_pattern_2: attribute axis of Split are not expected in past path")
127            return None
128        past = split.input[0]
129
130        past_k_nodes = self.model.match_parent_path(concat_k, ["Squeeze", "Split"], [0, 0])
131        if past_k_nodes is None:
132            logger.debug("match_past_pattern_2: failed to match past_k_nodes path")
133            return None
134        past_k = past_k_nodes[-1].input[0]
135
136        if past != past_k:
137            logger.info("match_past_pattern_2: expect past to be same")
138            return None
139
140        return past
141
142    def match_present(self, concat_v, input_name_to_nodes):
143        unsqueeze_present_v = self.model.find_first_child_by_type(
144            concat_v, "Unsqueeze", input_name_to_nodes, recursive=False
145        )
146        if not unsqueeze_present_v:
147            logger.info("expect unsqueeze for present")
148            return None
149        concat_present = self.model.find_first_child_by_type(
150            unsqueeze_present_v, "Concat", input_name_to_nodes, recursive=False
151        )
152        if not concat_present:
153            logger.info("expect concat for present")
154            return None
155
156        present = concat_present.output[0]
157        return present
158
159    def cast_attention_mask(self, input_name):
160        if input_name in self.casted_attention_mask:
161            attention_mask_input_name = self.casted_attention_mask[input_name]
162        elif self.model.find_graph_input(input_name):
163            casted, attention_mask_input_name = self.utils.cast_graph_input_to_int32(input_name)
164            self.casted_attention_mask[input_name] = attention_mask_input_name
165        else:
166            attention_mask_input_name, cast_node = self.utils.cast_input_to_int32(input_name)
167            self.casted_attention_mask[input_name] = attention_mask_input_name
168        return attention_mask_input_name
169
170
171class FusionGptAttention(FusionGptAttentionPastBase):
172    """
173    Fuse GPT-2 Attention with past state subgraph into one Attention node.
174    """
175
176    def __init__(self, model: OnnxModel, num_heads: int):
177        super().__init__(model, num_heads)
178
179    def create_attention_node(
180        self,
181        fc_weight,
182        fc_bias,
183        gemm_qkv,
184        past,
185        present,
186        input,
187        output,
188        mask,
189        is_unidirectional,
190    ):
191        attention_node_name = self.model.create_node_name("GptAttention")
192        attention_node = helper.make_node(
193            "Attention",
194            inputs=[input, fc_weight, fc_bias, mask, past],
195            outputs=[attention_node_name + "_output", present],
196            name=attention_node_name,
197        )
198        attention_node.domain = "com.microsoft"
199        attention_node.attribute.extend(
200            [
201                helper.make_attribute("num_heads", self.num_heads),
202                helper.make_attribute("unidirectional", 1 if is_unidirectional else 0),
203            ]
204        )
205
206        if self.mask_filter_value is not None:
207            attention_node.attribute.extend([helper.make_attribute("mask_filter_value", float(self.mask_filter_value))])
208
209        matmul_node = helper.make_node(
210            "MatMul",
211            inputs=[attention_node_name + "_output", gemm_qkv.input[1]],
212            outputs=[attention_node_name + "_matmul_output"],
213            name=attention_node_name + "_matmul",
214        )
215
216        add_node = helper.make_node(
217            "Add",
218            inputs=[attention_node_name + "_matmul_output", gemm_qkv.input[2]],
219            outputs=[output],
220            name=attention_node_name + "_add",
221        )
222        self.nodes_to_add.extend([attention_node, matmul_node, add_node])
223        self.node_name_to_graph_name[attention_node.name] = self.this_graph_name
224        self.node_name_to_graph_name[matmul_node.name] = self.this_graph_name
225        self.node_name_to_graph_name[add_node.name] = self.this_graph_name
226
227    def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
228        past = None
229        present = None
230        return_indice = []
231
232        is_normalize_node_skiplayernorm = normalize_node.op_type == "SkipLayerNormalization"
233        qkv_nodes = None
234
235        if not is_normalize_node_skiplayernorm:
236            qkv_nodes = self.model.match_parent_path(
237                normalize_node,
238                ["Add", "Reshape", "Gemm", "Reshape", "Reshape", "Transpose", "MatMul"],
239                [0, None, 0, 0, 0, 0, 0],
240                output_name_to_node=output_name_to_node,
241                return_indice=return_indice,
242            )
243        else:
244            qkv_nodes = self.model.match_parent_path(
245                normalize_node,
246                ["Reshape", "Gemm", "Reshape", "Reshape", "Transpose", "MatMul"],
247                [None, 0, 0, 0, 0, 0],
248                output_name_to_node=output_name_to_node,
249                return_indice=return_indice,
250            )
251
252        if qkv_nodes is None:
253            return
254
255        another_input = None
256        if not is_normalize_node_skiplayernorm:
257            (
258                add_qkv,
259                reshape_qkv,
260                gemm_qkv,
261                reshape_1,
262                reshape_2,
263                transpose_qkv,
264                matmul_qkv,
265            ) = qkv_nodes
266
267            another_input = add_qkv.input[1 - return_indice[0]]
268        else:
269            (
270                reshape_qkv,
271                gemm_qkv,
272                reshape_1,
273                reshape_2,
274                transpose_qkv,
275                matmul_qkv,
276            ) = qkv_nodes
277
278        v_nodes = self.model.match_parent_path(matmul_qkv, ["Concat", "Transpose", "Reshape", "Split"], [1, 1, 0, 0])
279        if v_nodes is None:
280            logger.debug("fuse_attention: failed to match v path")
281            return
282        (concat_v, transpose_v, reshape_v, split_fc) = v_nodes
283
284        # Try match pattern using Gemm + LayerNormalization
285        fc_nodes = self.model.match_parent_path(
286            split_fc,
287            ["Reshape", "Gemm", "Reshape", "LayerNormalization"],
288            [0, 0, 0, 0],
289            output_name_to_node,
290        )
291
292        # Try match pattern using Gemm + SkipLayerNormalization
293        if fc_nodes is None:
294            fc_nodes = self.model.match_parent_path(
295                split_fc,
296                ["Reshape", "Gemm", "Reshape", "SkipLayerNormalization"],
297                [0, 0, 0, 0],
298                output_name_to_node,
299            )
300
301        # Try match pattern using MatMul
302        if fc_nodes is None:
303            # LayerNormalization
304            fc_nodes = self.model.match_parent_path(
305                split_fc,
306                ["Add", "MatMul", "LayerNormalization"],
307                [0, None, 0],
308                output_name_to_node,
309            )
310
311            # SkipLayerNormalization
312            if fc_nodes is None:
313                fc_nodes = self.model.match_parent_path(
314                    split_fc,
315                    ["Add", "MatMul", "SkipLayerNormalization"],
316                    [0, None, 0],
317                    output_name_to_node,
318                )
319
320            if fc_nodes is None:
321                logger.debug("fuse_attention: failed to match fc path")
322                return
323
324            fc_weight = fc_nodes[1].input[1]
325            i, _ = self.model.get_constant_input(fc_nodes[0])
326            fc_bias = fc_nodes[0].input[i]
327        else:
328            fc_weight = fc_nodes[1].input[1]
329            fc_bias = fc_nodes[1].input[2]
330
331        layernorm_before_attention = fc_nodes[-1]
332
333        # `another_input` will be non-None only if
334        # (1) SkipLayerNorm fusion wasn't turned ON
335        # (2) SkipLayerNorm fusion was turned ON but upstream layer's LayerNorm + Add was not
336        # fused into a SkipLayerNorm. This can happen if the shapes to the Add node are different.
337        # So, keep the following check if SkipLayerNorm fusion is turned ON or OFF.
338        if another_input is not None and another_input not in layernorm_before_attention.input:
339            logger.debug("Upstream Add and (Skip)LayerNormalization shall have one same input")
340            return
341
342        is_unidirectional = True
343        slice_mask = None
344        input_mask_nodes = None
345        concat_k_to_match = None
346        qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Sub", "Mul", "Div", "MatMul"], [0, 0, 0, 0, 0])
347        if qk_nodes is not None:
348            (softmax_qk, sub_qk, mul_qk, div_qk, matmul_qk) = qk_nodes
349            mask_nodes = self.model.match_parent_path(
350                sub_qk,
351                [
352                    "Mul",
353                    "Sub",
354                    "Slice",
355                    "Slice",
356                    "Unsqueeze",
357                    "Sub",
358                    "Squeeze",
359                    "Slice",
360                    "Shape",
361                    "Div",
362                ],
363                [1, 0, 1, 0, 1, 0, 0, 0, 0, 0],
364            )
365            if mask_nodes is None:
366                logger.debug("fuse_attention: failed to match unidirectional mask path")
367                return
368            div_mask = mask_nodes[-1]
369            slice_mask = mask_nodes[3]
370
371            if div_qk != div_mask:
372                logger.debug("fuse_attention: skip since div_qk != div_mask")
373                return
374
375            if len(mask_nodes) > 1 and mask_nodes[0].op_type == "Mul":
376                _, mul_val = self.model.get_constant_input(mask_nodes[0])
377                if mul_val != -10000:
378                    self.mask_filter_value = -mul_val
379
380        else:
381            # New pattern for gpt2 from PyTorch 1.5.0 and Transformers 2.9.0.
382            i, qk_nodes, _ = self.model.match_parent_paths(
383                matmul_qkv,
384                [
385                    (["Softmax", "Where", "Div", "MatMul"], [0, 0, 1, 0]),
386                    (["Softmax", "Add", "Where", "Div", "MatMul"], [0, 0, None, 1, 0]),
387                ],
388                output_name_to_node,
389            )
390            if qk_nodes is None:
391                logger.debug("fuse_attention: failed to match qk nodes")
392                return
393
394            where_qk = qk_nodes[-3]
395            div_qk = qk_nodes[-2]
396            matmul_qk = qk_nodes[-1]
397
398            if i == 1:
399                add_qk = qk_nodes[1]
400                _, input_mask_nodes, _ = self.model.match_parent_paths(
401                    add_qk,
402                    [
403                        (
404                            ["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze", "Reshape"],
405                            [None, 0, 1, 0, 0, 0],
406                        ),
407                        (
408                            ["Mul", "Sub", "Unsqueeze", "Unsqueeze", "Reshape"],
409                            [None, 0, 1, 0, 0],
410                        ),
411                        (
412                            ["Mul", "Sub", "Unsqueeze", "Unsqueeze"],
413                            [None, 0, 1, 0],
414                        ),  # useless cast and reshape are removed.
415                    ],
416                    output_name_to_node,
417                )
418                if input_mask_nodes is None:
419                    logger.debug("fuse_attention: failed to match input attention mask path")
420                    return
421                if len(input_mask_nodes) > 1 and input_mask_nodes[0].op_type == "Mul":
422                    _, mul_val = self.model.get_constant_input(input_mask_nodes[0])
423                    if mul_val != -10000:
424                        self.mask_filter_value = mul_val
425
426            i, mask_nodes, _ = self.model.match_parent_paths(
427                where_qk,
428                [
429                    (
430                        ["Cast", "Slice", "Slice", "Unsqueeze", "Sub", "Squeeze", "Slice", "Shape"],
431                        [0, 0, 0, 1, 0, 0, 0, 0],
432                    ),
433                    # For Transformers >= 4.27, causal mask uses torch.bool instead of torch.uint8, so no Cast to bool.
434                    (
435                        ["Slice", "Slice", "Unsqueeze", "Sub", "Squeeze", "Slice", "Shape"],
436                        [0, 0, 1, 0, 0, 0, 0],
437                    ),
438                ],
439                output_name_to_node,
440            )
441            if mask_nodes is None:
442                # TODO: match mask path for GPT2LMHeadModel_BeamSearchStep.
443                logger.debug("fuse_attention: failed to match mask path")
444                return
445
446            slice_mask = mask_nodes[2 if i == 0 else 1]
447
448            div_or_concat = self.model.get_parent(mask_nodes[-1], 0, output_name_to_node)
449            if div_or_concat.op_type == "Div":
450                div_mask = div_or_concat
451                if div_qk != div_mask:
452                    logger.debug("fuse_attention: skip since div_qk != div_mask")
453                    return
454            elif div_or_concat.op_type == "Concat":
455                concat_k_to_match = div_or_concat
456            else:
457                logger.debug("fuse_attention: failed to match mask path")
458
459        # Validate that the mask data is either lower triangular (unidirectional) or all ones
460        mask_data = self.model.get_constant_value(slice_mask.input[0])
461        if not (
462            isinstance(mask_data, np.ndarray)
463            and len(mask_data.shape) == 4
464            and mask_data.shape[:2] == (1, 1)
465            and mask_data.shape[2] == mask_data.shape[3]
466        ):
467            logger.debug("fuse_attention: skip since mask shape is not 1x1xWxW")
468            return
469
470        if np.allclose(mask_data, np.ones_like(mask_data)):
471            is_unidirectional = False
472        elif not np.allclose(mask_data, np.tril(np.ones_like(mask_data))):
473            logger.debug("fuse_attention: skip since mask is neither lower triangular nor ones")
474            return
475
476        q_nodes = self.model.match_parent_path(matmul_qk, ["Transpose", "Reshape", "Split"], [0, 0, 0])
477        if q_nodes is None:
478            logger.debug("fuse_attention: failed to match q path")
479            return
480        (transpose_q, reshape_q, split_q) = q_nodes
481        if split_fc != split_q:
482            logger.debug("fuse_attention: skip since split_fc != split_q")
483            return
484
485        k_nodes = self.model.match_parent_path(matmul_qk, ["Concat", "Transpose", "Reshape", "Split"], [1, 1, 0, 0])
486        if k_nodes is None:
487            # This pattern is from pytorch 1.7.1 and transformers 4.6.1
488            k_nodes = self.model.match_parent_path(
489                matmul_qk,
490                ["Transpose", "Concat", "Transpose", "Reshape", "Split"],
491                [1, 0, 1, 0, 0],
492            )
493            if k_nodes is None:
494                logger.debug("fuse_attention: failed to match k path")
495                return
496            else:
497                (_, concat_k, transpose_k, reshape_k, split_k) = k_nodes
498        else:
499            (concat_k, transpose_k, reshape_k, split_k) = k_nodes
500        if split_fc != split_k:
501            logger.debug("fuse_attention: skip since split_fc != split_k")
502            return
503
504        if concat_k_to_match and concat_k != concat_k_to_match:
505            logger.debug("fuse_attention: skip since concat_k != concat_k_to_match")
506            return
507
508        attention_mask_input_name = ""
509        if input_mask_nodes is not None:
510            input_name = input_mask_nodes[-1].input[0]
511            attention_mask_input_name = self.cast_attention_mask(input_name)
512
513        # Match past and present paths
514        past = self.match_past_pattern_1(concat_k, concat_v, output_name_to_node) or self.match_past_pattern_2(
515            concat_k, concat_v, output_name_to_node
516        )
517        if past is None:
518            logger.info("fuse_attention: failed to match past path")
519            return
520        if not self.model.find_graph_input(past):
521            logger.debug("past is not graph input.")
522            # For GPT2LMHeadModel_BeamSearchStep, there is an extra Gather node to select beam index so it is not graph input.
523
524        present = self.match_present(concat_v, input_name_to_nodes)
525        if present is None:
526            logger.info("fuse_attention: failed to match present path")
527            return
528        if not self.model.find_graph_output(present):
529            logger.info("expect present to be graph output")
530            return
531
532        self.create_attention_node(
533            fc_weight,
534            fc_bias,
535            gemm_qkv,
536            past,
537            present,
538            layernorm_before_attention.output[0],
539            reshape_qkv.output[0],
540            attention_mask_input_name,
541            is_unidirectional,
542        )
543
544        # we rely on prune_graph() to clean old subgraph nodes:
545        # qk_nodes + q_nodes + k_nodes + v_nodes + mask_nodes + [reshape_qkv, transpose_qkv, matmul_qkv]
546        self.prune_graph = True
547 
codekingpro/portable-devtools · Team Ai