o
    ;ήc)A                     @   s  d dl Z d dlZd dlZd dlmZmZ d dlZd dlmZ d dl	m
Z
mZ d dlmZ d dlmZ d dlmZ d dlmZ d d	lmZ d d
lmZ eeZeddfeddfedd fedd feddfedd feddfdZdddg fdededee dee def
ddZ		 	 	d.de
dedededee f
dd Z 		 	 				d/d!ededededee dee ded"efd#d$Z!dedeeef fd%d&Z"d'd( Z#d)d* Z$d+d, Z%ed-kre%  dS dS )0    N)DictOptional)FusionOptions)
ModelProto
load_model)BartOnnxModel)BertOnnxModel)BertOnnxModelKeras)BertOnnxModelTF)Gpt2OnnxModel)TnlrOnnxModelpytorch   tf2onnx
keras2onnx)bartbertbert_tf
bert_kerasgpt2gpt2_tftnlrFc   onnx_model_pathuse_gpuoptimized_model_path	opt_levelreturnc           
      C   s&  |dv sJ ddl }|rd| vrtd | S | }|dkr'|jj|_n|dkr1|jj|_n|jj	|_|du rK| dd }d	
|||rHd
nd}||_i }|rV||d< |sf|j| |fddgi|}	n|j| |fddgi|}	d|	 v s{J tj|rtj|sJ td
| |S )a  
    Use onnxruntime to optimize model.

    Args:
        onnx_model_path (str): the path of input onnx model.
        use_gpu (bool): whether the optimized model is targeted to run in GPU.
        optimized_model_path (str or None): the path of optimized model.
        opt_level (int): graph optimization level.
        disabled_optimizers (List[str]): a list of names of disabled optimizers
    Returns:
        optimized_model_path (str): the path of optimized model
    )r      r   r   NCUDAExecutionProviderz3There is no gpu for onnxruntime to do optimization.r   r   z{}_o{}_{}.onnxgpucpudisabled_optimizers	providersCPUExecutionProviderz)Save optimized model by onnxruntime to {})onnxruntimeget_available_providersloggererrorSessionOptionsGraphOptimizationLevelORT_ENABLE_BASICgraph_optimization_levelORT_ENABLE_EXTENDEDORT_ENABLE_ALLformatoptimized_model_filepathInferenceSessionget_providersospathexistsisfiledebug)
r   r   r   r   r#   r&   sess_optionspath_prefixkwargssession r=   I/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/optimizer.pyoptimize_by_onnxruntime5   sJ   

r?   r   model
model_type	num_headshidden_sizeoptimization_optionsc           
      C   s   |dkr|dks|dkrt d t| \}}}| jr-|| jkr-t d| d| j d |du r5t|}|| ||}|| |  d|j_dd	lm	}	 |	|j_
|S )
ad  Optimize Model by graph fusion logic.

    Note that ONNXRuntime graph optimizations (like constant folding) will not be applied. So it is better to enable
    constant folding during exporting ONNX model, or run optimize_by_onnxruntime on the model first like optimize_model.

    For BERT model, num_heads and hidden_size are optional. For other model types, you need specify these parameters.

    Args:
        model (ModelProto): model object
        model_type (str, optional): model type - like bert, bert_tf, bert_keras or gpt2. Defaults to 'bert'.
        num_heads (int, optional): number of attention heads. Defaults to 0.
                                   0 allows detect the parameter from graph automatically (for model_type "bert" only).
        hidden_size (int, optional): hidden size. Defaults to 0.
                                     0 allows detect the parameter from graph automatically (for model_type "bert" only).
        optimization_options (FusionOptions, optional): optimization options that turn on/off some fusions. Defaults to None.

     Returns:
        object of an optimizer class.
    r   r   TPlease specify parameters of num_heads and hidden_size when model_type is not 'bert'z&Model producer not matched: Expected "z", Got "z0".Please specify correct --model_type parameter.Nzonnxruntime.transformers)__version__)r(   warningMODEL_TYPESproducer_namer   optimizetopological_sortr@   r&   rF   producer_version)
r@   rA   rB   rC   rD   optimizer_classproducer_	optimizeronnxruntime_versionr=   r=   r>   optimize_by_fusionp   s    

rR   inputonly_onnxruntimec                 C   s   |du s
|dv s
J |dkr|dks|dkrt d t| \}}	}
|du r(|
}d}|dkr?|r2g ng d}t| |||d}n|dkrJt| d	dd
}|rS|sSt d t|pW| }|rb||||}nt|||||}|ryt| t d	| |S )a	  Optimize Model by OnnxRuntime and/or python fusion logic.

    ONNX Runtime has graph optimizations (https://onnxruntime.ai/docs/resources/graph-optimizations.html).
    However, the coverage is limited. We also have graph fusions that implemented in Python to improve the coverage.
    They can combined: ONNX Runtime will run first when opt_level > 0, then graph fusions in Python will be applied.

    To use ONNX Runtime only and no Python fusion logic, use only_onnxruntime flag and a positive opt_level like
        optimize_model(input, opt_level=1, use_gpu=False, only_onnxruntime=True)

    When opt_level is None, we will choose default optimization level according to model type.

    When opt_level is 0 and only_onnxruntime is False, only python fusion logic is used and onnxruntime is disabled.

    When opt_level > 1, use_gpu shall set properly since the optimized graph might contain operators for GPU or CPU only.
    If your model is intended for GPU inference only (especially float16 or mixed precision model), it is recommended to
    set use_gpu to be True, otherwise the model is not optimized for GPU inference.

    For BERT model, num_heads and hidden_size are optional. For other model types, you need specify these parameters.

    Args:
        input (str): input model path.
        model_type (str, optional): model type - like bert, bert_tf, bert_keras or gpt2. Defaults to 'bert'.
        num_heads (int, optional): number of attention heads. Defaults to 0.
                                   0 allows detect the parameter from graph automatically (for model_type "bert" only).
        hidden_size (int, optional): hidden size. Defaults to 0.
                                     0 allows detect the parameter from graph automatically (for model_type "bert" only).
        optimization_options (FusionOptions, optional): optimization options that turn on/off some fusions. Defaults to None.
        opt_level (int, optional): onnxruntime graph optimization level (0, 1, 2 or 99) or None. Defaults to None.
                                   When the value is None, default value (1 for bert and gpt2, 0 for other model types) will be used.
                                   When the level > 0, onnxruntime will be used to optimize model first.
        use_gpu (bool, optional): use gpu or not for onnxruntime. Defaults to False.
        only_onnxruntime (bool, optional): only use onnxruntime to optimize model, and no python fusion. Defaults to False.

     Returns:
        object of an optimizer class.
    Nr   r   r   r   r   r   rE   r   )MatMulScaleFusionMatMulAddFusionSimplifiedLayerNormFusionGemmActivationFusionBiasSoftmaxFusion)r   r   r#   F)r   r   zKPlease specify a positive value for opt_level when only_onnxruntime is TruezRemove temporary model: {})
r(   rG   rH   r?   r   rR   r4   remover8   r0   )rS   rA   rB   rC   rD   r   r   rT   rM   	_producerdefault_opt_leveltemp_model_pathr#   r@   rP   r=   r=   r>   optimize_model   s<   .


r_   c                 C   s   t | ddd}t|}| S )z
    Get counter of fused operators in optimized model.

    Args:
        optimized_model_path (str): the path of onnx model.

    Returns:
        A dictionary with operator type as key, and count as value
    NT)r0   load_external_data)r   r   get_fused_operator_statistics)r   r@   rP   r=   r=   r>   get_fusion_statistics	  s   
rb   c                  C   sj  t jdd} | jddtdd | jddtdd | jd	d
tjdtt ddt  d | jdd
t	ddd | jdd
t	ddd | jdd
ddd | j
d
d | jdd
ddd | j
d
d t|  | jdd
ddd | j
d
d | jd d
dd!d | j
d
d" | jd#d
dd$d | j
d
d% | jd&d
t	g d'd d(d) | jd*d
dd+d | j
d
d, |  }|S )-NzuGraph optimization tool for ONNX Runtime. It transforms ONNX graph to use optimized operators for Transformer models.)descriptionz--inputTzinput onnx model path)requiredtypehelpz--outputzoptimized onnx model pathz--model_typeFr   z!Model type selected in the list: z, )rd   re   defaultchoicesrf   z--num_headsr   znumber of attention heads like 12 for bert-base and 16 for bert-large. Default is 0 to detect automatically for BERT. For other model type, this parameter need specify correctly.)rd   re   rg   rf   z--hidden_sizezhidden size like 768 for bert-base and 1024 for bert-large. Default is 0 to detect automatically for BERT. For other model type, this parameter need specify correctly.z--input_int32
store_truezyUse int32 (instead of int64) inputs. It could avoid unnecessary data cast when EmbedLayerNormalization is fused for BERT.)rd   actionrf   )input_int32z	--float16zConvert all weights and nodes in float32 to float16. It has potential loss in precision compared to mixed precision conversion (see convert_float_to_float16).)float16z	--verbosezshow debug information.verbosez	--use_gpuzZUse GPU for inference. Set this flag if your model is intended for GPU when opt_level > 1.)r   z--only_onnxruntimez<optimized by onnxruntime only, and no graph fusion in Python)rT   z--opt_levelrU   zonnxruntime optimization level. 0 will disable onnxruntime graph optimization. The recommended value is 1. When opt_level > 1 is used, optimized model for GPU might not run in CPU. Level 2 and 99 are intended for --only_onnxruntime.)rd   re   rh   rg   rf   z--use_external_data_formatz4use external data format to store large model (>2GB))use_external_data_format)argparseArgumentParseradd_argumentstrlowerlistrH   keysjoinintset_defaultsr   add_arguments
parse_args)parserargsr=   r=   r>   _parse_arguments  s   
	
	r~   c                 C   s&   | rt jddd d S t jdd d S )NDEBUGz8[%(filename)s:%(lineno)s - %(funcName)20s()] %(message)s)levelfmtz%(funcName)20s: %(message)s)r   )coloredlogsinstallrm   r=   r=   r>   _setup_loggert  s   
r   c               
   C   s   t  } t| j td|   tj| jtj| j	kr#t
d t| }t| j| j| j| j| j|| j| jd}| jrD|jdd | jrK|  || j	| j | r^td d S td d S )Nz
arguments:zYSpecified the same input and output path. Note that this may overwrite the original model)r   rD   r   rT   T)keep_io_typesz#The model has been fully optimized.zThe model has been optimized.)r~   r   rn   r(   r8   r4   r5   realpathrS   outputrG   r   parser_   rA   rB   rC   r   r   rT   rl   convert_float_to_float16rk   change_graph_inputs_to_int32save_model_to_filero   is_fully_optimizedinfo)r}   rD   rP   r=   r=   r>   main~  s0   


r   __main__)r   r   r   N)r   r   r   NNFF)&rp   loggingr4   typingr   r   r   fusion_optionsr   onnxr   r   onnx_model_bartr   onnx_model_bertr   onnx_model_bert_kerasr	   onnx_model_bert_tfr
   onnx_model_gpt2r   onnx_model_tnlrr   	getLogger__name__r(   rH   rs   boolrx   r?   rR   r_   rb   r~   r   r   r=   r=   r=   r>   <module>   s   

=
8
c\
%
