Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_base.py142 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from collections import defaultdict
6from collections.abc import Sequence
7from logging import getLogger
8from typing import Any
9
10import numpy as np
11from onnx import NodeProto, TensorProto, helper
12from onnx_model import OnnxModel
13
14logger = getLogger(__name__)
15
16
17class Fusion:
18    """
19    Base class for Graph Fusion
20    """
21
22    def __init__(
23        self,
24        model: OnnxModel,
25        fused_op_type: str,
26        search_op_types: str | list[str],
27        description: str = "",
28    ):
29        self.search_op_types: list[str] = [search_op_types] if isinstance(search_op_types, str) else search_op_types
30        self.fused_op_type: str = fused_op_type
31        self.description: str = f"{fused_op_type}({description})" if description else fused_op_type
32        self.model: OnnxModel = model
33        self.nodes_to_remove: list = []
34        self.nodes_to_add: list = []
35        self.prune_graph: bool = False
36        self.node_name_to_graph_name: dict = {}
37        self.this_graph_name: str | None = None
38        # It is optional that subclass updates fused_count since we will also check nodes_to_add to get counter.
39        self.fused_count: defaultdict = defaultdict(int)
40
41    def increase_counter(self, fused_op_name: str):
42        """
43        Increase counter of a fused operator.
44        """
45        self.fused_count[fused_op_name] += 1
46
47    def fuse(
48        self,
49        node: NodeProto,
50        input_name_to_nodes: dict[str, list[NodeProto]],
51        output_name_to_node: dict[str, NodeProto],
52    ):
53        """Interface for fusion that starts from a node"""
54        raise NotImplementedError
55
56    def apply(self):
57        """
58        Apply graph fusion on the whole model graph.
59        It searched nodes of given operators, and start fusion on each of those nodes.
60        """
61        logger.debug(f"start {self.description} fusion...")
62        input_name_to_nodes = self.model.input_name_to_nodes()
63        output_name_to_node = self.model.output_name_to_node()
64
65        # This assumes that two search ops will not be fused at same time!
66        for search_op_type in self.search_op_types:
67            for node in self.model.get_nodes_by_op_type(search_op_type):
68                graph = self.model.get_graph_by_node(node)
69                if graph is None:
70                    raise Exception("Can not find node in any graph")
71                self.this_graph_name = graph.name
72                self.fuse(node, input_name_to_nodes, output_name_to_node)
73
74        op_list = [node.op_type for node in self.nodes_to_add]
75        if self.fused_count:
76            for key, value in self.fused_count.items():
77                if value:
78                    logger.info(f"Fused {key}: {value}")
79        else:
80            count = op_list.count(self.fused_op_type)
81            if count > 0:
82                logger.info(f"Fused {self.description}: {count}")
83
84        self.model.remove_nodes(self.nodes_to_remove)
85        self.model.add_nodes(self.nodes_to_add, self.node_name_to_graph_name)
86
87        if self.prune_graph:
88            self.model.prune_graph()
89        elif self.nodes_to_remove or self.nodes_to_add:
90            self.model.update_graph()
91
92    def add_initializer(self, name: str, data_type: int, dims: Sequence[int], vals: Any, raw: bool = True):
93        if raw:
94            if not isinstance(vals, np.ndarray):
95                np_type = helper.tensor_dtype_to_np_dtype(data_type)
96                bytes = np.array(vals, dtype=np_type).tobytes()
97            else:
98                bytes = vals.tobytes()
99            tensor = helper.make_tensor(
100                name=name,
101                data_type=data_type,
102                dims=dims,
103                vals=bytes,
104                raw=True,
105            )
106        else:
107            tensor = helper.make_tensor(
108                name=name,
109                data_type=data_type,
110                dims=dims,
111                vals=vals,
112                raw=False,
113            )
114
115        self.model.add_initializer(tensor, self.this_graph_name)
116        return tensor
117
118    def remove_initializer(self, tensor: TensorProto):
119        self.model.remove_initializer(tensor)
120
121    def add_nodes_to_remove(self, nodes: list[NodeProto]):
122        # Some nodes are shared between paths (e.g. rotary embedding nodes in the Q and K paths).
123        # When path A is fused, its shared nodes are added to `self.nodes_to_remove`. But when path B
124        # is fused, its shared nodes are also added to `self.nodes_to_remove`. When the nodes are
125        # iteratively removed from `self.nodes_to_remove`, path A's shared nodes are removed first.
126        # Since path A's shared nodes are removed, path B's shared nodes are not removed because they
127        # were previously removed for path A. This causes an error to print in remove_node that a node
128        # has failed to be removed.
129        #
130        # To avoid this error, we pre-emptively check if the shared nodes are already in `self.nodes_to_remove`.
131        # We could alternatively convert `self.nodes_to_remove` to a set to avoid this issue, but there could
132        # be scenarios where the nodes need to be removed in a specific order and converting to a set would
133        # lose this order.
134        for node in nodes:
135            if node not in self.nodes_to_remove:
136                self.nodes_to_remove.append(node)
137
138    def add_nodes_to_remove_with_nodes_to_keep(self, nodes: list[NodeProto], nodes_to_keep: list[NodeProto]):
139        for node in nodes:
140            if node not in self.nodes_to_remove and node not in nodes_to_keep:
141                self.nodes_to_remove.append(node)
142 
codekingpro/portable-devtools · Team Ai