Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_exporter.py720 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.  See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6
7import logging
8import os
9from pathlib import Path
10
11import numpy
12import torch
13from affinity_helper import AffinitySetting
14from benchmark_helper import OptimizerInfo, Precision, create_onnxruntime_session
15from huggingface_models import MODEL_CLASSES
16from quantize_helper import QuantizeHelper
17from torch_onnx_export_helper import torch_onnx_export
18from transformers import AutoConfig, AutoFeatureExtractor, AutoTokenizer, LxmertConfig, TransfoXLConfig
19
20from onnxruntime.transformers.models.gpt2.gpt2_helper import (
21    PRETRAINED_GPT2_MODELS,
22    GPT2ModelNoPastState,
23    TFGPT2ModelNoPastState,
24)
25
26os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2"
27
28logger = logging.getLogger(__name__)
29
30# Workaround by replacing torch.triu using self-defined op
31# Since torch.triu cannot be exported to ONNX. See https://github.com/pytorch/pytorch/issues/32968
32torch_func = {"triu": torch.triu}
33
34
35def triu_onnx(x, diagonal=0, out=None):
36    assert out is None
37    assert len(x.shape) == 2 and x.size(0) == x.size(1)
38
39    torch_triu = torch_func["triu"]
40    template = torch_triu(torch.ones((1024, 1024), dtype=torch.uint8), diagonal)
41    mask = template[: x.size(0), : x.size(1)]
42    return torch.where(mask.bool(), x, torch.zeros_like(x))
43
44
45def replace_torch_functions():
46    torch.triu = triu_onnx
47
48
49def restore_torch_functions():
50    torch.triu = torch_func["triu"]
51
52
53def create_onnxruntime_input(vocab_size, batch_size, sequence_length, input_names, config, data_type=numpy.int64):
54    if config.model_type in ["vit", "swin"]:
55        input_ids = numpy.random.rand(batch_size, 3, config.image_size, config.image_size).astype(numpy.float32)
56        inputs = {"pixel_values": input_ids}
57        return inputs
58
59    input_ids = numpy.random.randint(low=0, high=vocab_size - 1, size=(batch_size, sequence_length), dtype=data_type)
60    inputs = {"input_ids": input_ids}
61
62    if "attention_mask" in input_names:
63        attention_mask = numpy.ones([batch_size, sequence_length], dtype=data_type)
64        inputs["attention_mask"] = attention_mask
65
66    if "token_type_ids" in input_names:
67        segment_ids = numpy.zeros([batch_size, sequence_length], dtype=data_type)
68        inputs["token_type_ids"] = segment_ids
69
70    if config.is_encoder_decoder:
71        inputs["decoder_input_ids"] = input_ids
72
73    if isinstance(config, LxmertConfig):
74        inputs["visual_feats"] = numpy.random.randn(1, 1, config.visual_feat_dim).astype(numpy.float32)
75        inputs["visual_pos"] = numpy.random.randn(1, 1, config.visual_pos_dim).astype(numpy.float32)
76    if isinstance(config, TransfoXLConfig):
77        inputs["tf_transfo_xl_model/transformer/pos_emb/einsum/Einsum/inputs_1:0"] = numpy.zeros(
78            [config.hidden_size], dtype=numpy.float32
79        )
80    return inputs
81
82
83def filter_inputs(inputs, input_names):
84    remaining_model_inputs = {}
85    for input_name in input_names:
86        if input_name in inputs:
87            remaining_model_inputs[input_name] = inputs[input_name]
88    return remaining_model_inputs
89
90
91def flatten(inputs):
92    return [[flatten(i) for i in inputs] if isinstance(inputs, (list, tuple)) else inputs]
93
94
95def update_flatten_list(inputs, res_list):
96    for i in inputs:
97        res_list.append(i) if not isinstance(i, (list, tuple)) else update_flatten_list(i, res_list)
98    return res_list
99
100
101def build_dynamic_axes(example_inputs, outputs_flatten):
102    sequence_length = example_inputs["input_ids"].shape[-1]
103
104    dynamic_axes = {key: {0: "batch_size", 1: "seq_len"} for key in example_inputs}
105
106    output_names = ["output_" + str(i + 1) for i in range(len(outputs_flatten))]
107    for i, output_name in enumerate(output_names):
108        dynamic_axes[output_name] = {0: "batch_size"}
109        dims = outputs_flatten[i].shape
110        for j, dim in enumerate(dims):
111            if dim == sequence_length:
112                dynamic_axes[output_name].update({j: "seq_len"})
113    return dynamic_axes, output_names
114
115
116def validate_onnx_model(
117    onnx_model_path,
118    example_inputs,
119    example_outputs_flatten,
120    use_gpu,
121    fp16,
122    output_names=None,
123):
124    test_session = create_onnxruntime_session(onnx_model_path, use_gpu, enable_all_optimization=False)
125    if test_session is None:
126        logger.error(f"{onnx_model_path} is an invalid ONNX model")
127        return False
128
129    logger.info(f"{onnx_model_path} is a valid ONNX model")
130
131    # Compare the inference result with PyTorch or Tensorflow
132    example_ort_inputs = {k: t.numpy() for k, t in example_inputs.items()}
133    example_ort_outputs = test_session.run(output_names, example_ort_inputs)
134    if len(example_outputs_flatten) != len(example_ort_outputs):
135        logger.error(
136            f"Number of output tensors expected {len(example_outputs_flatten)}, got {len(example_ort_outputs)}"
137        )
138        return False
139
140    for i in range(len(example_outputs_flatten)):
141        abs_diff = numpy.amax(numpy.abs(example_ort_outputs[i] - example_outputs_flatten[i].cpu().numpy()))
142        if abs_diff > 1e-4:
143            logger.info(f"Max absolute diff={abs_diff} for output tensor {i}")
144
145        rtol = 5e-02 if fp16 else 1e-4
146        atol = 1e-01 if fp16 else 1e-4
147        if not numpy.allclose(
148            example_ort_outputs[i],
149            example_outputs_flatten[i].cpu().numpy(),
150            rtol=rtol,
151            atol=atol,
152        ):
153            logger.error(f"Output tensor {i} is not close: rtol={rtol}, atol={atol}")
154            return False
155
156    logger.info(f"inference result of onnxruntime is validated on {onnx_model_path}")
157    return True
158
159
160def get_onnx_file_path(
161    onnx_dir: str,
162    model_name: str,
163    input_count: int,
164    optimized_by_script: bool,
165    use_gpu: bool,
166    precision: Precision,
167    optimized_by_onnxruntime: bool,
168    use_external_data: bool,
169):
170    from re import sub  # noqa: PLC0415
171
172    normalized_model_name = sub(r"[^a-zA-Z0-9_]", "_", model_name)
173
174    if not optimized_by_script:
175        filename = f"{normalized_model_name}_{input_count}"
176    else:
177        device = "gpu" if use_gpu else "cpu"
178        filename = f"{normalized_model_name}_{input_count}_{precision}_{device}"
179
180    if optimized_by_onnxruntime:
181        filename += "_ort"
182
183    directory = onnx_dir
184    # ONNXRuntime will not write external data so the raw and optimized models shall be in same directory.
185    if use_external_data and not optimized_by_onnxruntime:
186        directory = os.path.join(onnx_dir, filename)
187        if not os.path.exists(directory):
188            os.makedirs(directory)
189
190    return os.path.join(directory, f"{filename}.onnx")
191
192
193def add_filename_suffix(file_path: str, suffix: str) -> str:
194    """
195    Append a suffix at the filename (before the extension).
196    Args:
197        path: pathlib.Path The actual path object we would like to add a suffix
198        suffix: The suffix to add
199    Returns: path with suffix appended at the end of the filename and before extension
200    """
201    path = Path(file_path)
202    return str(path.parent.joinpath(path.stem + suffix).with_suffix(path.suffix))
203
204
205def optimize_onnx_model_by_ort(onnx_model_path, ort_model_path, use_gpu, overwrite, model_fusion_statistics):
206    if overwrite or not os.path.exists(ort_model_path):
207        Path(ort_model_path).parent.mkdir(parents=True, exist_ok=True)
208        from optimizer import get_fusion_statistics, optimize_by_onnxruntime  # noqa: PLC0415
209
210        # Use onnxruntime to optimize model, which will be saved to *_ort.onnx
211        _ = optimize_by_onnxruntime(
212            onnx_model_path,
213            use_gpu=use_gpu,
214            optimized_model_path=ort_model_path,
215            opt_level=99,
216        )
217        model_fusion_statistics[ort_model_path] = get_fusion_statistics(ort_model_path)
218    else:
219        logger.info(f"Skip optimization since model existed: {ort_model_path}")
220
221
222def optimize_onnx_model(
223    onnx_model_path,
224    optimized_model_path,
225    model_type,
226    num_attention_heads,
227    hidden_size,
228    use_gpu,
229    precision,
230    use_raw_attention_mask,
231    overwrite,
232    model_fusion_statistics,
233    use_external_data_format,
234    optimization_options=None,
235):
236    if overwrite or not os.path.exists(optimized_model_path):
237        Path(optimized_model_path).parent.mkdir(parents=True, exist_ok=True)
238
239        from fusion_options import FusionOptions  # noqa: PLC0415
240        from optimizer import optimize_model  # noqa: PLC0415
241
242        if optimization_options is None:
243            optimization_options = FusionOptions(model_type)
244        optimization_options.use_raw_attention_mask(use_raw_attention_mask)
245        if precision == Precision.FLOAT16:
246            optimization_options.enable_gelu_approximation = True
247        if precision == Precision.INT8:
248            optimization_options.enable_embed_layer_norm = False
249
250        # For swin models, the num_attention_heads is a list, which isn't supported yet, so set to 0 for now
251        if model_type == "swin":
252            num_attention_heads = 0
253            hidden_size = 0
254
255        # Use script to optimize model.
256        # Use opt_level <= 1 for models to be converted to fp16, because some fused op (like FusedGemm) has only fp32 and no fp16.
257        # It is better to be conservative so we use opt_level=0 here, in case MemcpyFromHost is added to the graph by OnnxRuntime.
258        opt_model = optimize_model(
259            onnx_model_path,
260            model_type,
261            num_heads=num_attention_heads,
262            hidden_size=hidden_size,
263            opt_level=0,
264            optimization_options=optimization_options,
265            use_gpu=use_gpu,
266            only_onnxruntime=False,
267        )
268        if model_type == "bert_keras" or model_type == "bert_tf":
269            opt_model.use_dynamic_axes()
270
271        model_fusion_statistics[optimized_model_path] = opt_model.get_fused_operator_statistics()
272
273        if precision == Precision.FLOAT16:
274            opt_model.convert_float_to_float16(keep_io_types=True)
275
276        opt_model.save_model_to_file(optimized_model_path, use_external_data_format)
277    else:
278        logger.info(f"Skip optimization since model existed: {optimized_model_path}")
279
280
281def modelclass_dispatcher(model_name, custom_model_class):
282    if custom_model_class is not None:
283        if custom_model_class in MODEL_CLASSES:
284            return custom_model_class
285        else:
286            raise Exception("Valid model class: " + " ".join(MODEL_CLASSES))
287
288    if model_name in PRETRAINED_GPT2_MODELS:
289        return "GPT2ModelNoPastState"
290
291    import re  # noqa: PLC0415
292
293    if re.search("-squad$", model_name) is not None:
294        return "AutoModelForQuestionAnswering"
295    elif re.search("-mprc$", model_name) is not None:
296        return "AutoModelForSequenceClassification"
297    elif re.search("gpt2", model_name) is not None:
298        return "AutoModelWithLMHead"
299
300    return "AutoModel"
301
302
303def load_pretrained_model(model_name, config, cache_dir, custom_model_class, is_tf_model=False):
304    model_class_name = modelclass_dispatcher(model_name, custom_model_class)
305
306    if model_class_name == "GPT2ModelNoPastState":
307        if is_tf_model:
308            return TFGPT2ModelNoPastState.from_pretrained(model_name, config=config, cache_dir=cache_dir)
309        else:
310            return GPT2ModelNoPastState.from_pretrained(model_name, config=config, cache_dir=cache_dir)
311
312    if is_tf_model:
313        model_class_name = "TF" + model_class_name
314
315    transformers_module = __import__("transformers", fromlist=[model_class_name])
316    logger.info(f"Model class name: {model_class_name}")
317    model_class = getattr(transformers_module, model_class_name)
318
319    return model_class.from_pretrained(model_name, config=config, cache_dir=cache_dir)
320
321
322def load_pt_model(model_name, model_class, cache_dir, config_modifier):
323    config = AutoConfig.from_pretrained(model_name, cache_dir=cache_dir)
324    if hasattr(config, "return_dict"):
325        config.return_dict = False
326
327    config_modifier.modify(config)
328
329    model = load_pretrained_model(model_name, config=config, cache_dir=cache_dir, custom_model_class=model_class)
330
331    return config, model
332
333
334def load_tf_model(model_name, model_class, cache_dir, config_modifier):
335    config = AutoConfig.from_pretrained(model_name, cache_dir=cache_dir)
336
337    config_modifier.modify(config)
338    # Loading tf model from transformers limits the cpu affinity to {0} when KMP_AFFINITY is set
339    # Restore the affinity after model loading for expected ORT performance
340    affinity_setting = AffinitySetting()
341    affinity_setting.get_affinity()
342    model = load_pretrained_model(
343        model_name,
344        config=config,
345        cache_dir=cache_dir,
346        custom_model_class=model_class,
347        is_tf_model=True,
348    )
349    affinity_setting.set_affinity()
350
351    return config, model
352
353
354# For test only
355def load_pt_model_from_tf(model_name):
356    # Note that we could get pt model from tf, but model source and its structure in this case is different from directly using
357    # load_pt_model() and load_tf_model() even with the same name. Therefore it should not be used for comparing with them
358    from convert_tf_models_to_pytorch import tf2pt_pipeline  # noqa: PLC0415
359
360    config, model = tf2pt_pipeline(model_name)
361
362    return config, model
363
364
365def validate_and_optimize_onnx(
366    model_name,
367    use_external_data_format,
368    model_type,
369    onnx_dir,
370    input_names,
371    use_gpu,
372    precision,
373    optimize_info,
374    validate_onnx,
375    use_raw_attention_mask,
376    overwrite,
377    config,
378    model_fusion_statistics,
379    onnx_model_path,
380    example_inputs,
381    example_outputs_flatten,
382    output_names,
383    fusion_options,
384):
385    is_valid_onnx_model = True
386    if validate_onnx:
387        is_valid_onnx_model = validate_onnx_model(
388            onnx_model_path,
389            example_inputs,
390            example_outputs_flatten,
391            use_gpu,
392            False,
393            output_names,
394        )
395    if optimize_info.name == OptimizerInfo.NOOPT.name:
396        return onnx_model_path, is_valid_onnx_model, config.vocab_size
397
398    if (
399        optimize_info.name == OptimizerInfo.BYSCRIPT.name
400        or precision == Precision.FLOAT16
401        or precision == Precision.INT8
402    ):  # Use script (optimizer.py) to optimize
403        optimized_model_path = get_onnx_file_path(
404            onnx_dir,
405            model_name,
406            len(input_names),
407            True,
408            use_gpu,
409            precision,
410            False,
411            use_external_data_format,
412        )
413        optimize_onnx_model(
414            onnx_model_path,
415            optimized_model_path,
416            model_type,
417            config.num_attention_heads,
418            config.hidden_size,
419            use_gpu,
420            precision,
421            use_raw_attention_mask,
422            overwrite,
423            model_fusion_statistics,
424            use_external_data_format,
425            fusion_options,
426        )
427
428        onnx_model_path = optimized_model_path
429        if validate_onnx:
430            is_valid_onnx_model = validate_onnx_model(
431                onnx_model_path,
432                example_inputs,
433                example_outputs_flatten,
434                use_gpu,
435                precision == Precision.FLOAT16,
436                output_names,
437            )
438
439        if precision == Precision.INT8:
440            logger.info(f"Quantizing model: {onnx_model_path}")
441            QuantizeHelper.quantize_onnx_model(onnx_model_path, onnx_model_path, use_external_data_format)
442            logger.info(f"Finished quantizing model: {onnx_model_path}")
443
444    if optimize_info.name == OptimizerInfo.BYORT.name:  # Use OnnxRuntime to optimize
445        if is_valid_onnx_model:
446            ort_model_path = add_filename_suffix(onnx_model_path, "_ort")
447            optimize_onnx_model_by_ort(
448                onnx_model_path,
449                ort_model_path,
450                use_gpu,
451                overwrite,
452                model_fusion_statistics,
453            )
454
455    return (
456        onnx_model_path,
457        is_valid_onnx_model,
458        config.num_labels if model_type in ["vit", "swin"] else config.vocab_size,
459    )
460
461
462def export_onnx_model_from_pt(
463    model_name,
464    opset_version,
465    use_external_data_format,
466    model_type,
467    model_class,
468    config_modifier,
469    cache_dir,
470    onnx_dir,
471    input_names,
472    use_gpu,
473    precision,
474    optimizer_info,
475    validate_onnx,
476    use_raw_attention_mask,
477    overwrite,
478    model_fusion_statistics,
479    fusion_options,
480):
481    config, model = load_pt_model(model_name, model_class, cache_dir, config_modifier)
482    # config, model = load_pt_model_from_tf(model_name)
483    model.cpu()
484
485    example_inputs = None
486    max_input_size = None
487
488    if model_type in ["vit", "swin"]:
489        image_processor = AutoFeatureExtractor.from_pretrained(model_name, cache_dir=cache_dir)
490        data = numpy.random.randint(
491            low=0, high=256, size=config.image_size * config.image_size * 3, dtype=numpy.uint8
492        ).reshape(config.image_size, config.image_size, 3)
493
494        example_inputs = image_processor(data, return_tensors="pt")
495    else:
496        tokenizer = AutoTokenizer.from_pretrained(model_name, cache_dir=cache_dir)
497        max_input_size = tokenizer.model_max_length
498        example_inputs = tokenizer.encode_plus("This is a sample input", return_tensors="pt")
499
500    example_inputs = filter_inputs(example_inputs, input_names)
501
502    example_outputs = model(**example_inputs)
503
504    assert isinstance(example_outputs, (list, tuple)), f"type of output is not list or tuple: {type(example_outputs)}"
505
506    # Flatten is needed for gpt2 and distilgpt2.
507    example_outputs_flatten = flatten(example_outputs)
508    example_outputs_flatten = update_flatten_list(example_outputs_flatten, [])
509
510    onnx_model_path = get_onnx_file_path(
511        onnx_dir,
512        model_name,
513        len(input_names),
514        False,
515        use_gpu,
516        precision,
517        False,
518        use_external_data_format,
519    )
520
521    if overwrite or not os.path.exists(onnx_model_path):
522        logger.info(f"Exporting ONNX model to {onnx_model_path}")
523        Path(onnx_model_path).parent.mkdir(parents=True, exist_ok=True)
524
525        dynamic_axes = None
526        output_names = None
527
528        if model_type in ["vit", "swin"]:
529            dynamic_axes, output_names = {key: {0: "pixel_values"} for key in example_inputs}, ["logits"]
530        else:
531            dynamic_axes, output_names = build_dynamic_axes(example_inputs, example_outputs_flatten)
532
533        replace_torch_functions()
534        torch_onnx_export(
535            model=model,
536            args=tuple(example_inputs.values()),
537            f=onnx_model_path,
538            input_names=list(example_inputs.keys()),
539            output_names=output_names,
540            dynamic_axes=dynamic_axes,
541            do_constant_folding=True,
542            opset_version=opset_version,
543            use_external_data_format=use_external_data_format,
544        )
545        restore_torch_functions()
546    else:
547        logger.info(f"Skip export since model existed: {onnx_model_path}")
548
549    onnx_model_file, is_valid_onnx_model, vocab_size = validate_and_optimize_onnx(
550        model_name,
551        use_external_data_format,
552        model_type,
553        onnx_dir,
554        input_names,
555        use_gpu,
556        precision,
557        optimizer_info,
558        validate_onnx,
559        use_raw_attention_mask,
560        overwrite,
561        config,
562        model_fusion_statistics,
563        onnx_model_path,
564        example_inputs,
565        example_outputs_flatten,
566        None,
567        fusion_options,
568    )
569
570    return onnx_model_file, is_valid_onnx_model, vocab_size, max_input_size
571
572
573def export_onnx_model_from_tf(
574    model_name,
575    opset_version,
576    use_external_data_format,
577    model_type,
578    model_class,
579    config_modifier,
580    cache_dir,
581    onnx_dir,
582    input_names,
583    use_gpu,
584    precision,
585    optimizer_info,
586    validate_onnx,
587    use_raw_attention_mask,
588    overwrite,
589    model_fusion_statistics,
590    fusion_options,
591):
592    # Use CPU to export
593    import tensorflow as tf  # noqa: PLC0415
594
595    tf.config.set_visible_devices([], "GPU")
596
597    tokenizer = AutoTokenizer.from_pretrained(model_name, cache_dir=cache_dir)
598    # Fix "Using pad_token, but it is not set yet" error.
599    if tokenizer.pad_token is None:
600        tokenizer.add_special_tokens({"pad_token": "[PAD]"})
601    max_input_size = tokenizer.model_max_length
602
603    config, model = load_tf_model(model_name, model_class, cache_dir, config_modifier)
604    model.resize_token_embeddings(len(tokenizer))
605
606    example_inputs = tokenizer.encode_plus(
607        "This is a sample input",
608        return_tensors="tf",
609        max_length=max_input_size,
610        padding="max_length",
611        truncation=True,
612    )
613    example_inputs = filter_inputs(example_inputs, input_names)
614
615    if config.is_encoder_decoder:
616        example_inputs["decoder_input_ids"] = tokenizer.encode_plus(
617            "This is a sample input",
618            return_tensors="tf",
619            max_length=max_input_size,
620            padding="max_length",
621            truncation=True,
622        ).input_ids
623    if model_name == "unc-nlp/lxmert-base-uncased":
624        example_inputs["visual_feats"] = tf.random.normal([1, 1, config.visual_feat_dim])
625        example_inputs["visual_pos"] = tf.random.normal([1, 1, config.visual_pos_dim])
626
627    try:
628        # Use no past state for these models
629        if config.use_cache:
630            config.use_cache = False
631    except Exception:
632        pass
633
634    example_outputs = model(example_inputs, training=False)
635    output_names = None
636
637    # For xlnet models, only compare the last_hidden_state output.
638    if model_name == "xlnet-base-cased" or model_name == "xlnet-large-cased":
639        output_names = ["last_hidden_state"]
640        example_outputs = example_outputs["last_hidden_state"]
641
642    # Flatten is needed for gpt2 and distilgpt2. Output name sorting is needed for tf2onnx outputs to match onnx outputs.
643    from tensorflow.python.util import nest  # noqa: PLC0415
644
645    example_outputs_flatten = nest.flatten(example_outputs)
646
647    onnx_model_path = get_onnx_file_path(
648        onnx_dir,
649        model_name,
650        len(input_names),
651        False,
652        use_gpu,
653        precision,
654        False,
655        use_external_data_format,
656    )
657    tf_internal_model_path = onnx_model_path[:-5] if use_external_data_format else onnx_model_path
658
659    if overwrite or not os.path.exists(tf_internal_model_path):
660        logger.info(f"Exporting ONNX model to {onnx_model_path}")
661        if not use_external_data_format:
662            Path(tf_internal_model_path).parent.mkdir(parents=True, exist_ok=True)
663
664        import zipfile  # noqa: PLC0415
665
666        import tf2onnx  # noqa: PLC0415
667
668        tf2onnx.logging.set_level(tf2onnx.logging.ERROR)
669        specs = []
670        for name, value in example_inputs.items():
671            dims = [None] * len(value.shape)
672            specs.append(tf.TensorSpec(tuple(dims), value.dtype, name=name))
673        _, _ = tf2onnx.convert.from_keras(
674            model,
675            input_signature=tuple(specs),
676            opset=opset_version,
677            large_model=use_external_data_format,
678            output_path=tf_internal_model_path,
679        )
680        if use_external_data_format:
681            # need to unpack the zip for run_onnxruntime()
682            with zipfile.ZipFile(tf_internal_model_path, "r") as z:
683                z.extractall(os.path.dirname(tf_internal_model_path))
684            tf_internal_model_path = os.path.join(os.path.dirname(tf_internal_model_path), "__MODEL_PROTO.onnx")
685            if os.path.exists(onnx_model_path):
686                os.remove(onnx_model_path)
687            os.rename(tf_internal_model_path, onnx_model_path)
688
689    else:
690        logger.info(f"Skip export since model existed: {onnx_model_path}")
691
692    model_type = model_type + "_tf"
693    optimized_onnx_path, is_valid_onnx_model, vocab_size = validate_and_optimize_onnx(
694        model_name,
695        use_external_data_format,
696        model_type,
697        onnx_dir,
698        input_names,
699        use_gpu,
700        precision,
701        optimizer_info,
702        validate_onnx,
703        use_raw_attention_mask,
704        overwrite,
705        config,
706        model_fusion_statistics,
707        onnx_model_path,
708        example_inputs,
709        example_outputs_flatten,
710        output_names,
711        fusion_options,
712    )
713
714    return (
715        optimized_onnx_path,
716        is_valid_onnx_model,
717        vocab_size,
718        max_input_size,
719    )
720 
codekingpro/portable-devtools · Team Ai