o
    ;ήcL!                     @   s   d dl Z d dlZd dlZd dlZd dlZd dlZd dlmZmZm	Z	 ej
ej
ej
edd d dlmZmZmZmZ edZdd Z								
ddedededededefddZdd Zedkrle  dS dS )    N)PRETRAINED_MT5_MODELSPRETRAINED_T5_MODELST5Helperz..)	Precisioncreate_onnxruntime_sessionprepare_environmentsetup_logger c               
   C   s  t  } tt }| jdddtd tdd| d | jddtd	d	d
gdd | jddttjdddd | jddttjdddd | jdddddd | j	dd | jddddd | j	dd | jdddt
t
jt
jt
jgd d | jd!ddd" | j	dd# | jd$d%ddd" | j	dd& | jd'd(ddd)d | j	dd* | jd+d,ddd-d | j	dd. | jd/ddd0d | j	dd1 | jd2ddd3d | j	dd4 | jd5ddd6d | j	dd7 |  }|S )8Nz-mz--model_name_or_pathFr   z2Model path, or pretrained model name in the list: z, )requireddefaulttypehelpz--model_typet5mt5z&Model type: either t5 (default) or mt5)r
   r   r   choicesr   z--cache_dir.cache_modelsz%Directory to cache pre-trained models)r
   r   r   r   z--outputonnx_modelszOutput directoryz-oz--optimize_onnx
store_truez'Use optimizer.py to optimize onnx model)r
   actionr   )optimize_onnxz	--use_gpuzuse GPU for inference)use_gpuz-pz--precisionzKPrecision of model to run. fp32 for full precision, fp16 for half precisionz	--verbose)r
   r   )verbosez-ez--use_external_data_format)use_external_data_formatz-sz--use_decoder_start_tokenz]Use config.decoder_start_token_id. Otherwise, add an extra graph input for decoder_input_ids.)use_decoder_start_tokenz-wz--overwritezoverwrite existing ONNX model)	overwritez--disable_auto_mixed_precisionz(use pure fp16 instead of mixed precision)disable_auto_mixed_precisionz#--separate_encoder_and_decoder_initzHDo not merge encode and decoder init. Output 3 instead of 2 onnx models.)!separate_encoder_and_decoder_initz--use_int64_inputszJUse int64 instead of int32 for input_ids, position_ids and attention_mask.)use_int64_inputs)argparseArgumentParserr   r   add_argumentstrjoinospathset_defaultsr   FLOAT32FLOAT16
parse_args)parserpretrained_modelsargs r-   Y/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/t5/convert_to_onnx.pyparse_arguments   s   		

r/   FTr   r   merge_encoder_and_decoder_initr   r   use_int32_inputs
model_typec              
   C   s  t |rdnd}t| |||	|}|d j}|s#|jdkr#td g }| D ]\}}|	| d| }tj
|| |dd}|
sGtj|setd	|  t|	|}tj|||||| |d
 ntd|  |st|tjkrtj
|| |d t| dd}|
stj|std|  tj|||tjk|j|j|| d ntd|  n|}t|||rddgndgd}t   t||||}W d    n1 sw   Y  td|  |dkrtd || q)|S )Nzcuda:0cpudecoder   z2Try use_external_data_format when model size > 2GB_F)suffix
new_folderzExporting ONNX model to )use_decoder_input_idsr1   z#Skip exporting: existed ONNX model zOptimizing model to )auto_mixed_precisionz$Skip optimizing: existed ONNX model CUDAExecutionProviderCPUExecutionProvider)r   providerz1PyTorch and OnnxRuntime results max difference = g-C6?z-PyTorch and OnnxRuntime results are NOT close)torchdevicer   
load_modelconfig
num_layersloggerinfoitemstoget_onnx_pathr$   r%   existscopydeepcopyexport_onnxr   r'   r"   r   r(   	num_headshidden_sizer   no_gradverify_onnxwarningappend)model_name_or_path	cache_dir
output_dirr   r   r   	precisionr   r   r0   r   r   r1   r2   r?   modelsrA   output_pathsnamemodelfilename_suffix	onnx_pathcloned_modeloutput_pathort_sessionmax_diffr-   r-   r.   export_onnx_models   sz   






r`   c                  C   s   t  } t| j td|   | j}| jds| jntj	
| j}t||| j | jtjkr7| js7J d| jtjkrD| jsDJ d| jrLtd t| j||| j| j| j| j| j| j| j | j| j| j | j}td|  d S )Nz
Arguments:z.onnxz"fp16/int8 requires --optimize_onnxzfp16 requires --use_gpuz1Graph optimization for T5 is not implemented yet.zDone! Outputs: )r/   r   r   rC   rD   rS   outputendswithr$   r%   dirnamer   r   rU   r   r'   r   r(   rP   r`   rR   r   r   r   r   r   r   r2   )r,   rS   rT   rW   r-   r-   r.   main   s:   
 
rd   __main__)FTFFTr   )r   rI   loggingr$   sysr>   	t5_helperr   r   r   r%   rQ   r#   rc   __file__benchmark_helperr   r   r   r   	getLoggerrC   r/   boolr"   r`   rd   __name__r-   r-   r-   r.   <module>   sD    
z	

^(
