Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
t5_helper.py316 linesDownload Raw Back to t5
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# -------------------------------------------------------------------------
5
6import logging
7import os
8from pathlib import Path
9
10import torch
11from float16 import float_to_float16_max_diff
12from onnx_model import OnnxModel
13from optimizer import optimize_model
14from t5_decoder import T5Decoder, T5DecoderHelper
15from t5_encoder_decoder_init import T5EncoderDecoderInit, T5EncoderDecoderInitHelper
16from transformers import MT5ForConditionalGeneration, T5ForConditionalGeneration
17
18from onnxruntime import InferenceSession
19
20logger = logging.getLogger(__name__)
21
22
23def _torch_load_weights_only(path: str, **kwargs):
24    try:
25        return torch.load(path, weights_only=True, **kwargs)
26    except TypeError:
27        logger.warning(
28            "Current PyTorch version does not support torch.load(..., weights_only=True); "
29            "falling back to default torch.load behavior for %s.",
30            path,
31        )
32        return torch.load(path, **kwargs)
33
34
35PRETRAINED_T5_MODELS = ["t5-small", "t5-base", "t5-large", "t5-3b", "t5-11b"]
36PRETRAINED_MT5_MODELS = [
37    "google/mt5-small",
38    "google/mt5-base",
39    "google/mt5-large",
40    "google/mt5-xl",
41    "google/mt5-xxl",
42]
43
44
45class T5Helper:
46    @staticmethod
47    def get_onnx_path(
48        output_dir: str,
49        model_name_or_path: str,
50        suffix: str = "",
51        new_folder: bool = False,
52    ) -> str:
53        """Build onnx path
54
55        Args:
56            output_dir (str): output directory
57            model_name_or_path (str): pretrained model name, or path to the model checkpoint
58            suffix (str, optional): suffix like "_encoder" or "_decoder_fp16" will be appended to file name. Defaults to None.
59            new_folder (bool, optional): create a new directory for the model. Defaults to False.
60
61        Returns:
62            str: path of onnx model
63        """
64        model_name = model_name_or_path
65        if os.path.isdir(model_name_or_path):
66            model_name = Path(model_name_or_path).parts[-1]
67        else:
68            model_name.split("/")[-1]
69
70        model_name += suffix
71
72        directory = os.path.join(output_dir, model_name) if new_folder else output_dir
73        return os.path.join(directory, model_name + ".onnx")
74
75    @staticmethod
76    def load_model(
77        model_name_or_path: str,
78        cache_dir: str,
79        device: torch.device,
80        model_type: str = "t5",
81        state_dict_path: str = "",
82        encoder_decoder_init: bool = False,
83    ) -> dict[str, T5EncoderDecoderInit | T5Decoder]:
84        """Load model given a pretrained name or path, then build models for ONNX conversion.
85
86        Args:
87            model_name_or_path (str): pretrained model name or path
88            cache_dir (str): cache directory
89            device (torch.device): device to run the model
90            model_type (str, optional): model type "t5" or "mt5"
91            state_dict_path(str, optional): state dictionary path
92            encoder_decoder_init (bool, optional): combine encoder and decoder kv cache initialization into one model.
93        Returns:
94            Dict[str, torch.nn.Module]: mapping from name to modules for ONNX conversion.
95        """
96        if model_type == "t5":
97            model = T5ForConditionalGeneration.from_pretrained(model_name_or_path, cache_dir=cache_dir)
98        elif model_type == "mt5":
99            model = MT5ForConditionalGeneration.from_pretrained(model_name_or_path, cache_dir=cache_dir)
100        else:
101            raise ValueError("only support mode_type=t5 or mt5")
102
103        if state_dict_path:
104            model.load_state_dict(_torch_load_weights_only(state_dict_path))
105
106        decoder = T5Decoder(model.decoder, model.lm_head, model.config)
107        decoder.eval().to(device)
108
109        encoder = T5EncoderDecoderInit(
110            model.encoder,
111            model.decoder,
112            model.lm_head,
113            model.config,
114            decoder_start_token_id=None,
115            output_cross_only=not encoder_decoder_init,
116        )
117
118        encoder_name = "encoder_decoder_init" if encoder_decoder_init else "encoder"
119        return {encoder_name: encoder, "decoder": decoder}
120
121    @staticmethod
122    def export_onnx(
123        model: T5Decoder | T5EncoderDecoderInit,
124        device: torch.device,
125        onnx_model_path: str,
126        verbose: bool = True,
127        use_external_data_format: bool = False,
128        use_decoder_input_ids: bool = True,
129        use_int32_inputs: bool = False,
130    ):
131        if isinstance(model, T5EncoderDecoderInit):
132            T5EncoderDecoderInitHelper.export_onnx(
133                model,
134                device,
135                onnx_model_path,
136                use_decoder_input_ids,
137                verbose,
138                use_external_data_format,
139                use_int32_inputs,
140            )
141        else:
142            T5DecoderHelper.export_onnx(
143                model,
144                device,
145                onnx_model_path,
146                verbose,
147                use_external_data_format,
148                use_int32_inputs,
149            )
150
151    @staticmethod
152    def auto_mixed_precision(
153        onnx_model: OnnxModel,
154        op_block_list: list[str] | None = None,
155        force_fp16_logits: bool = False,
156        use_symbolic_shape_infer: bool = True,
157    ):
158        """Convert model to mixed precision.
159           It detects whether original model has fp16 precision weights, and set parameters for float16 conversion automatically.
160        Args:
161            onnx_model (OnnxModel): optimized ONNX model
162            op_block_list (List[str], optional): operators need to run in fp32.
163            force_fp16_logits (bool, optional): force logits and last MatMul node to be in float16. Defaults to False.
164            use_symbolic_shape_infer (bool, optional): use symbolic shape inference to convert float to float16. Defaults to True.
165        Returns:
166            parameters(dict): a dictionary of parameters used in float16 conversion
167        """
168        if op_block_list is None:
169            op_block_list = [
170                "SimplifiedLayerNormalization",
171                "SkipSimplifiedLayerNormalization",
172                "Relu",
173                "Add",
174            ]
175
176        op_full_set = {node.op_type for node in onnx_model.nodes()}
177        fp32_op_set = set(op_block_list)
178        fp16_op_set = op_full_set.difference(fp32_op_set)
179        logger.info(f"fp32 op: {fp32_op_set} fp16 op: {fp16_op_set}")
180
181        # logits is the first output
182        logits_output_name = onnx_model.graph().output[0].name
183
184        # We use the weight in last MatMul node to detect whether the model is stored with float16 weights from training.
185        is_weight_fp16_precision = False
186        output_name_to_node = onnx_model.output_name_to_node()
187        assert logits_output_name in output_name_to_node
188        node = output_name_to_node[logits_output_name]
189        last_matmul_node = None
190        if node.op_type == "MatMul":
191            last_matmul_node = node
192            logger.info(f"Found last MatMul node for logits: {node.name}")
193            initializer = None
194            for input in node.input:
195                initializer = onnx_model.get_initializer(input)
196                if initializer is not None:
197                    break
198
199            # when the max difference of value after converting float to float16 is lower than a threshold (1e-6),
200            # we can deduce that the weights are stored in float16 precision.
201            max_diff = float_to_float16_max_diff(initializer)
202            logger.debug(f"max diff of converting weights in last MatMul node {node.name}: {max_diff}")
203            is_weight_fp16_precision = max_diff < 1e-6
204        else:
205            logger.warning(f"Failed to find MatMul node for logits. Found {node.op_type} of node {node.name}")
206
207        keep_io_types = []
208        node_block_list = []
209        if (not is_weight_fp16_precision) and (last_matmul_node is not None) and not force_fp16_logits:
210            # When original weight is float32 precision, keep logits and last MatMul in float32 could get better precision.
211            keep_io_types = [logits_output_name]
212            node_block_list = [last_matmul_node.name]
213
214        if "Add" not in op_block_list:
215            input_name_to_nodes = onnx_model.input_name_to_nodes()
216            fp32_add = 0
217            changed = True
218            add_nodes = onnx_model.get_nodes_by_op_type("Add")
219            while changed:
220                changed = False
221                for node in add_nodes:
222                    if node.name not in node_block_list:
223                        parents = onnx_model.get_parents(node, output_name_to_node)
224                        children = onnx_model.get_children(node, input_name_to_nodes)
225                        blocked_children = [
226                            child for child in children if child.op_type in op_block_list or child in node_block_list
227                        ]
228                        blocked_parents = [
229                            parent for parent in parents if parent.op_type in op_block_list or parent in node_block_list
230                        ]
231                        # If any child or parent is in fp32, we place the Add node to fp32.
232                        if (len(blocked_children) + len(blocked_parents)) > 0:
233                            node_block_list.append(node.name)
234                            fp32_add += 1
235                            changed = True
236            fp16_add = len(add_nodes) - fp32_add
237            logger.info(f"node counter of Add operator: fp32={fp32_add} fp16={fp16_add}")
238
239        logger.info(f"node_block_list: {node_block_list}")
240
241        parameters = {
242            "keep_io_types": keep_io_types,
243            "op_block_list": op_block_list,
244            "node_block_list": node_block_list,
245            "force_fp16_initializers": is_weight_fp16_precision,
246        }
247
248        logger.info(f"auto_mixed_precision parameters: {parameters}")
249        if use_symbolic_shape_infer:
250            onnx_model.convert_float_to_float16(use_symbolic_shape_infer=True, **parameters)
251        else:
252            # Workaround when symbolic shape inference fails.
253            # Need enable shape_infer_before_optimization in convert_to_onnx.py as well.
254            from float16 import convert_float_to_float16  # noqa: PLC0415
255
256            convert_float_to_float16(
257                onnx_model.model,
258                disable_shape_infer=True,
259                **parameters,
260            )
261
262        return parameters
263
264    @staticmethod
265    def optimize_onnx(
266        onnx_model_path: str,
267        optimized_model_path: str,
268        is_float16: bool,
269        num_attention_heads: int,
270        hidden_size: int,
271        use_external_data_format: bool = False,
272        auto_mixed_precision: bool = True,
273        use_gpu: bool = False,
274        force_fp16_io: bool = False,
275    ):
276        """Optimize ONNX model with an option to convert it to use mixed precision."""
277
278        from fusion_options import FusionOptions  # noqa: PLC0415
279
280        optimization_options = None
281        if is_float16:
282            optimization_options = FusionOptions("t5")
283            # SkipLayerNormalization is faster but might bring accuracy drop since it uses fp16 accumulation.
284            optimization_options.enable_skip_layer_norm = not auto_mixed_precision
285
286        m = optimize_model(
287            onnx_model_path,
288            model_type="t5",
289            num_heads=num_attention_heads,
290            hidden_size=hidden_size,
291            opt_level=0,
292            optimization_options=optimization_options,
293            use_gpu=use_gpu,
294        )
295
296        if is_float16:
297            if auto_mixed_precision:
298                T5Helper.auto_mixed_precision(m, force_fp16_logits=force_fp16_io)
299            else:
300                m.convert_model_float32_to_float16(cast_input_output=force_fp16_io)
301
302        m.save_model_to_file(optimized_model_path, use_external_data_format, all_tensors_to_one_file=True)
303
304    @staticmethod
305    def verify_onnx(
306        model: T5Decoder | T5EncoderDecoderInit,
307        ort_session: InferenceSession,
308        device: torch.device,
309        use_int32_inputs: bool,
310    ):
311        """Compare the result from PyTorch and OnnxRuntime to verify the ONNX model is good."""
312        if isinstance(model, T5EncoderDecoderInit):
313            return T5EncoderDecoderInitHelper.verify_onnx(model, ort_session, device, use_int32_inputs)
314
315        return T5DecoderHelper.verify_onnx(model, ort_session, device, use_int32_inputs)
316 
codekingpro/portable-devtools · Team Ai