codekingpro/portable-devtools
114k
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 