o
    ;ήcZ[                     @   s  d dl Z d dlZd dlZd dlmZ d dlZd dlZd dlmZ d dl	m
Z
mZmZ d dlmZ d dlmZ d dlmZ d dlmZmZmZmZ ejejejed	d
 d dlmZmZm Z  dej!d< e "e#Z$dej%iZ&dEddZ'dd Z(dd Z)ej*fddZ+dd Z,dd Z-dd Z.dd Z/	dFdd Z0d!e1d"e1d#e2d$e3d%e3d&ed'e3d(e3fd)d*Z4d+e1d,e1d-e1fd.d/Z5d0d1 Z6	dFd2d3Z7d4d5 Z8dGd7d8Z9d9d: Z:d;d< Z;d=d> Z<d?d@ Z=dAdB Z>dCdD Z?dS )H    N)Path)AffinitySetting)OptimizerInfo	Precisioncreate_onnxruntime_session)MODEL_CLASSES)QuantizeHelper)torch_onnx_export)
AutoConfigAutoTokenizerLxmertConfigTransfoXLConfigmodelsgpt2)PRETRAINED_GPT2_MODELSGPT2ModelNoPastStateTFGPT2ModelNoPastState2TF_CPP_MIN_LOG_LEVELtriuc                 C   s   |d u sJ t | jdkr| d| dksJ td }|tjdtjd|}|d | dd | df }t| | t	| S )N   r      r   )   r   dtype)
lenshapesize
torch_functorchonesuint8wherebool
zeros_like)xdiagonalout
torch_triutemplatemask r+   M/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/onnx_exporter.py	triu_onnx!   s   & r-   c                   C   s
   t t_d S N)r-   r   r   r+   r+   r+   r,   replace_torch_functions+   s   
r/   c                   C   s   t d t_d S )Nr   )r   r   r   r+   r+   r+   r,   restore_torch_functions/   s   r0   c           
      C   s   t jjd| d ||f|d}d|i}d|v r#t j||g|d}||d< d|v r4t j||g|d}	|	|d< |jr;||d< t|tr^t jdd|j	
t j|d	< t jdd|j
t j|d
< t|trot j|jgt jd|d< |S )Nr   r   )lowhighr   r   	input_idsattention_maskr   token_type_idsdecoder_input_idsvisual_feats
visual_posz@tf_transfo_xl_model/transformer/pos_emb/einsum/Einsum/inputs_1:0)numpyrandomrandintr    zerosis_encoder_decoder
isinstancer   randnvisual_feat_dimastypefloat32visual_pos_dimr   hidden_size)

vocab_size
batch_sizesequence_lengthinput_namesconfig	data_typer3   inputsr4   segment_idsr+   r+   r,   create_onnxruntime_input3   s$   



rM   c                 C   s&   i }|D ]}|| v r| | ||< q|S r.   r+   )rK   rH   remaining_model_inputs
input_namer+   r+   r,   filter_inputsM   s   rP   c                 C   s$   t | ttfrdd | D gS | gS )Nc                 S   s   g | ]}t |qS r+   )flatten.0ir+   r+   r,   
<listcomp>V   s    zflatten.<locals>.<listcomp>)r>   listtuple)rK   r+   r+   r,   rQ   U   s   $rQ   c                 C   s0   | D ]}t |ttfs||nt|| q|S r.   )r>   rV   rW   appendupdate_flatten_list)rK   res_listrT   r+   r+   r,   rY   Y   s    rY   c           
      C   s   | d j d }dd |  D }dd tt|D }t|D ]%\}}ddi||< || j }t|D ]\}}	|	|krC|| |d	i q2q||fS )
Nr3   c                 S   s   i | ]}|d ddqS )rF   seq_len)r   r   r+   )rS   keyr+   r+   r,   
<dictcomp>b   s    z&build_dynamic_axes.<locals>.<dictcomp>c                 S   s   g | ]
}d t |d  qS )output_r   )strrR   r+   r+   r,   rU   d   s    z&build_dynamic_axes.<locals>.<listcomp>r   rF   r\   )r   keysranger   	enumerateupdate)
example_inputsoutputs_flattenrG   dynamic_axesoutput_namesrT   output_namedimsjdimr+   r+   r,   build_dynamic_axes_   s   
rm   c              	   C   sN  t | |dd}|d u rt|  d dS t|  d dd | D }|||}t|t|krEtdt| dt|  dS tt|D ]Q}	t	t
||	 ||	    }
|
d	krntd
|
 d|	  |rrdnd	}|rxdnd	}tj||	 ||	   ||dstd|	 d| d|   dS qKtd|   dS )NF)enable_all_optimizationz is an invalid ONNX modelz is a valid ONNX modelc                 S   s   i | ]	\}}||  qS r+   )r9   )rS   ktr+   r+   r,   r^   ~   s    z'validate_onnx_model.<locals>.<dictcomp>z"Number of output tensors expected z, got g-C6?zMax absolute diff=z for output tensor g?g?)rtolatolzOutput tensor z is not close: rtol=z, atol=z0inference result of onnxruntime is validated on T)r   loggererrorinfoitemsrunr   rb   r9   amaxabscpuallclose)onnx_model_pathre   example_outputs_flattenuse_gpufp16rh   test_sessionexample_ort_inputsexample_ort_outputsrT   abs_diffrq   rr   r+   r+   r,   validate_onnx_modeln   s:   $	r   onnx_dir
model_nameinput_countoptimized_by_scriptr~   	precisionoptimized_by_onnxruntimeuse_external_datac                 C   s   ddl m} |dd|}	|s|	 d| }
n|rdnd}|	 d| d| d| }
|r/|
d7 }
| }|rG|sGtj| |
}tj|sGt| tj||
 dS )	Nr   )subz[^a-zA-Z0-9_]_gpurz   _ortz.onnx)rer   ospathjoinexistsmakedirs)r   r   r   r   r~   r   r   r   r   normalized_model_namefilenamedevice	directoryr+   r+   r,   get_onnx_file_path   s   

r   	file_pathsuffixreturnc                 C   s&   t | }t|j|j| |jS )a  
    Append a suffix at the filename (before the extension).
    Args:
        path: pathlib.Path The actual path object we would like to add a suffix
        suffix: The suffix to add
    Returns: path with suffix appended at the end of the filename and before extension
    )r   r`   parentjoinpathstemwith_suffixr   )r   r   r   r+   r+   r,   add_filename_suffix   s   r   c                 C   sh   |st j|s*t|jjddd ddlm}m} || ||dd}||||< d S t	
d|  d S )NTparentsexist_okr   )get_fusion_statisticsoptimize_by_onnxruntimec   )r~   optimized_model_path	opt_level'Skip optimization since model existed: )r   r   r   r   r   mkdir	optimizerr   r   rs   ru   )r|   ort_model_pathr~   	overwritemodel_fusion_statisticsr   r   r   r+   r+   r,   optimize_onnx_model_by_ort   s   r   c              
   C   s   |st j|slt|jjddd ddlm} ddlm	} |d u r&||}|
| tj|kr3d|_tj|kr;d|_|| |||d||dd}|dksO|d	krS|  | |	|< tj|krd|jdd
 |||
 d S td|  d S )NTr   r   )FusionOptions)optimize_modelF)	num_headsrD   r   optimization_optionsr~   only_onnxruntime
bert_kerasbert_tf)keep_io_typesr   )r   r   r   r   r   r   fusion_optionsr   r   r   use_raw_attention_maskr   FLOAT16enable_gelu_approximationINT8enable_embed_layer_normuse_dynamic_axesget_fused_operator_statisticsconvert_float_to_float16save_model_to_filers   ru   )r|   r   
model_typenum_attention_headsrD   r~   r   r   r   r   use_external_data_formatr   r   r   	opt_modelr+   r+   r,   optimize_onnx_model   s8   




r   c                 C   sz   |d kr|t v r
|S tddt  | tv rdS dd l}|d| d kr'dS |d| d kr1dS |d	| d kr;d
S dS )NzValid model class:  r   r   z-squad$AutoModelForQuestionAnsweringz-mprc$"AutoModelForSequenceClassificationr   AutoModelWithLMHead	AutoModel)r   	Exceptionr   r   r   search)r   custom_model_classr   r+   r+   r,   modelclass_dispatcher  s   r   Fc                 C   sz   t | |}|dkr|rtj| ||dS tj| ||dS |r!d| }td|gd}td|  t||}|j| ||dS )Nr   )rI   	cache_dirTFtransformers)fromlistzModel class name: )r   r   from_pretrainedr   
__import__rs   ru   getattr)r   rI   r   r   is_tf_modelmodel_class_nametransformers_modulemodel_classr+   r+   r,   load_pretrained_model$  s   

r   c                 C   s@   t j| |d}t|drd|_|| t| |||d}||fS )Nr   return_dictF)rI   r   r   )r
   r   hasattrr   modifyr   )r   r   r   config_modifierrI   modelr+   r+   r,   load_pt_model7  s   

r   c                 C   sH   t j| |d}|| t }|  t| |||dd}|  ||fS )Nr   T)rI   r   r   r   )r
   r   r   r   get_affinityr   set_affinity)r   r   r   r   rI   affinity_settingr   r+   r+   r,   load_tf_modelC  s   
r   c                 C   s    ddl m} || \}}||fS )Nr   )tf2pt_pipeline)convert_tf_models_to_pytorchr   )r   r   rI   r   r+   r+   r,   load_pt_model_from_tfX  s   r   c                 C   s  d}|rt ||||d|}|tjkr|||jfS |tjks'|tjks'|tjkrqt|| t	|d||d|}t
||||j|j|||	|
||| |}|rUt |||||tjk|}|tjkrqtd|  t||| td|  |tjkr|rt|d}t||||
| |||jfS )NTFzQuantizing model: zFinished quantizing model: r   )r   r   NOOPTrE   BYSCRIPTr   r   r   r   r   r   r   rD   rs   ru   r   quantize_onnx_modelBYORTr   r   )r   r   r   r   rH   r~   r   optimize_infovalidate_onnxr   r   rI   r   r|   re   r}   rh   r   is_valid_onnx_modelr   r   r+   r+   r,   validate_and_optimize_onnxb  s   


	

r   c                 C   sx  t | |||\}}|  tj| |d}| |jv r|j|  nd}|jddd}t||}|di |}t|tt	fsCJ dt
| t|}t|g }t|| t|d|	|
d|}|satj|std| t|jjd	d	d
 t||\}}t  t|t	| |t| ||d	||d	 t  ntd|  t| |||||	|
|||||||||d |\}}}||||fS )Nr   r   This is a sample inputpt)return_tensorsz%type of output is not list or tuple: FExporting ONNX model to {}Tr   )	r   argsfrH   rh   rg   do_constant_foldingopset_versionr   !Skip export since model existed: r+   )r   rz   r   r   max_model_input_sizesencode_plusrP   r>   rV   rW   typerQ   rY   r   r   r   r   r   rs   ru   formatr   r   r   rm   r/   r	   valuesra   r0   r   )r   r   r   r   r   r   r   r   rH   r~   r   optimizer_infor   r   r   r   r   rI   r   	tokenizermax_input_sizere   example_outputsr}   r|   rg   rh   onnx_model_filer   rE   r+   r+   r,   export_onnx_model_from_pt  sx   
 



r  c           (      C   s  dd l }|jg d tj| |d}|jd u r|ddi | |jv r(|j|  nd}t| |||\}}|	t
| |jdd|d	d
d}t||}|jrY|jdd|d	d
dj|d< | dkru|jdd|jg|d< |jdd|jg|d< z|jr|d|_W n   Y ||dd}d }| dks| dkrdg}|d }ddlm} ||}t|| t
|d|	|
d|}|r|d d n|}|stj|sYtd| |st|jj d
d
d dd l!}dd l"}|j#$|j#j% g }|& D ]\} }!d gt
|!j' }"|(|j)t*|"|!j+| d q|j,j-|t*||||d\}#}#|rX|.|d}$|$/tj0| W d    n	1 s6w   Y  tj1tj0|d}tj|rRt2| t3|| ntd|  |d }t4| |||||	|
|||||||||||\}%}&}'|%|&|'|fS ) Nr   GPUr   	pad_tokenz[PAD]r   r   tf
max_lengthT)r   r  padding
truncationr6   zunc-nlp/lxmert-base-uncasedr   r7   r8   F)trainingzxlnet-base-casedzxlnet-large-casedlast_hidden_state)nestr   r   )name)input_signatureopsetlarge_modeloutput_pathrz__MODEL_PROTO.onnxr   _tf)5
tensorflowrI   set_visible_devicesr   r   r	  add_special_tokensr   r   resize_token_embeddingsr   r   rP   r=   r3   r:   normalr@   rC   	use_cachetensorflow.python.utilr  rQ   r   r   r   r   rs   ru   r   r   r   r   zipfiletf2onnxlogging	set_levelERRORrv   r   rX   
TensorSpecrW   r   convert
from_kerasZipFile
extractalldirnamer   removerenamer   )(r   r   r   r   r   r   r   r   rH   r~   r   r  r   r   r   r   r   r
  r  r  rI   r   re   r  rh   r  r}   r|   tf_internal_model_pathr   r!  specsr  valuerj   r   zoptimized_onnx_pathr   rE   r+   r+   r,   export_onnx_model_from_tf  s   






r2  )r   Nr.   )F)@r"  r   syspathlibr   r9   r   affinity_helperr   benchmark_helperr   r   r   huggingface_modelsr   quantize_helperr   torch_onnx_export_helperr	   r   r
   r   r   r   r   rX   r   r*  __file__gpt2_helperr   r   r   environ	getLogger__name__rs   r   r   r-   r/   r0   int64rM   rP   rQ   rY   rm   r   r`   intr#   r   r   r   r   r   r   r   r   r   r   r  r2  r+   r+   r+   r,   <module>   sp    





,
!
6

[b