o
    ;ήc.                     @   s  d dl Z d dl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
Z
d dlZd dlZd dlmZ d dlmZ d dlmZmZ d dlmZmZ d dlmZ ejejejed	d	 d d
lmZ d dl m!Z! e "e#Z$G dd dej%j&Z'G dd dZ(G dd dZ)dS )    N)Path)ListOptionalUnion)PastKeyValuesHelper)T5DecoderInit)	T5EncoderT5EncoderInputs)	MT5ConfigT5Config)InferenceSessionz..)	OnnxModel)torch_onnx_exportc                       sr   e Zd ZdZ	ddejjdejjdejjdeee	f de
e f
 fdd	Z	dd
ejdejdejfddZ  ZS )T5EncoderDecoderInitz-A combination of T5Encoder and T5DecoderInit.Nencoderdecoderlm_headconfigdecoder_start_token_idc                    s0   t    || _t||| _t||||| _d S N)super__init__r   r   
t5_encoderr   t5_decoder_init)selfr   r   r   r   r   	__class__ a/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/t5/t5_encoder_decoder_init.pyr   "   s   
zT5EncoderDecoderInit.__init__encoder_input_idsencoder_attention_maskdecoder_input_idsc                 C   s,   |  ||}| |||\}}}||||fS r   )r   r   )r   r   r    r!   encoder_hidden_states	lm_logits	past_self
past_crossr   r   r   forward/   s
   
zT5EncoderDecoderInit.forwardr   )__name__
__module____qualname____doc__torchnnModuler   r   r
   r   intr   Tensorr&   __classcell__r   r   r   r   r      s,    
r   c                   @   sX   e Zd ZdddZe	ddeeef dededed	e	j
d
efddZdefddZdS )T5EncoderDecoderInitInputsNc                 C   s   || _ || _|| _d S r   )r   r    r!   )r   r   r    r!   r   r   r   r   =   s   
z#T5EncoderDecoderInitInputs.__init__Fr   
batch_sizeencode_sequence_lengthuse_decoder_input_idsdeviceuse_int32_inputsc           	      C   sX   t j||| j||d}d }|r$|rtjntj}tj|df||d| j }t|j	|j
|S )N)r6      )dtyper5   )r	   create_dummy
vocab_sizer+   int32int64onesr   r1   	input_idsattention_mask)	r   r2   r3   r4   r5   r6   encoder_inputsr!   r8   r   r   r   r9   B   s   	z'T5EncoderDecoderInitInputs.create_dummyreturnc                 C   s&   | j | jg}| jd ur|| j |S r   )r   r    r!   append)r   
input_listr   r   r   to_listY   s   
z"T5EncoderDecoderInitInputs.to_listr   )F)r'   r(   r)   r   staticmethodr   r   r
   r.   r+   r5   boolr9   r   rD   r   r   r   r   r1   <   s$    

r1   c                   @   s|   e Zd Ze				ddedejdedededed	efd
dZ	ede
fddZe	ddededejd	edef
ddZdS )T5EncoderDecoderInitHelperTFmodelr5   onnx_model_pathr4   verboseuse_external_data_formatr6   c                 C   s  t | tsJ tj| jdd|||d}| }tj| jjdd}	ddg|	 }
dd	g}d
}t	| jj
}t	| jj}t	| jj}dddddddd|dd|dd}|r`|d d|d|d< |	D ]}d|v rrd|d|d||< qbd|||d||< qbt c}tj|d}t|jjddd t| t||d||
|dd||d t|} | jjD ]%}|jjjjD ]}| dr|j!||||fv rt"|j!}|#  ||_$qqt%j&| ||dd W d   dS 1 sw   Y  dS )a  Export decoder to ONNX

        Args:
            model (T5EncoderDecoderInit): the model to export
            device (torch.device): device of decoder object
            onnx_model_path (str): onnx path
            verbose (bool, optional): print verbose information. Defaults to True.
            use_external_data_format (bool, optional): use external data format or not. Defaults to False.
              )r2   r3   r4   r5   r6   T)presentlogitsr"   r   r    1r2   r3   )r   r7   )r   r7   rL   )r   r    r"   rO   r!   cross)r   r7   rL   rM   zencoder_decoder_init.onnx)parentsexist_ok   )
argsfexport_paramsinput_namesoutput_namesdynamic_axesopset_versiondo_constant_foldingrK   rJ   	dim_param)save_as_external_dataall_tensors_to_one_fileN)'
isinstancer   r1   r9   r   rD   r   get_past_names
num_layersstr	num_headsd_modeld_kvrB   tempfileTemporaryDirectoryospathjoinr   parentmkdirr   tupleonnxloadgraphoutputtypetensor_typeshapedimHasFieldr]   r.   Clear	dim_valuer   save)rH   r5   rI   r4   rJ   rK   r6   inputsrC   present_namesrY   rX   sequence_lengthrd   hidden_size	head_sizerZ   nametmp_dir_nametemp_onnx_model_pathtensor	dim_protory   r   r   r   export_onnxa   s   

	


"z&T5EncoderDecoderInitHelper.export_onnxr{   c                 C   sf   t d t|j  t|j  d}|jdur+t|j  |d< | d|}|S )zRun inference of ONNX model.zstart onnxruntime_inference)r   r    Nr!   )	loggerdebugnumpyascontiguousarrayr   cpur    r!   run)ort_sessionr{   
ort_inputsort_outputsr   r   r   onnxruntime_inference   s   

z0T5EncoderDecoderInitHelper.onnxruntime_inference   r   	max_casesc                 C   s  |  }t|dk}g d}g }|d| D ]\}	}
tj| j|	|
|||d}t||}| }| | }|d  	 j
|d j
ksDJ t	t	|d  	 |d  }td|  |}|d  	 j
|d j
kspJ t	t	|d  	 |d  }td|  t||}td	| jj D ]#}t	t	|d	 |  	 |d	|   }td
| d|  qtd	| jj D ].}t	t	|d |  	 |d	d	| jj  |   }td| d|  t||}q|| td|	 d|
 d|  qt|S )zQCompare the result from PyTorch and OnnxRuntime to verify the ONNX model is good.rM   ))r      )r7   rL   )rM   r7   )      N)r4   r5   r6   r   zlogits max_diff=r7   zencoder_hidden_states max_diff=rL   zself attention past state z
 max_diff=zcross attention past state zbatch_size=z encode_sequence_length=z, max_diff=)
get_inputslenr1   r9   r   rG   r   rD   r   r   ru   amaxabsr   r   maxrangerb   rB   info)rH   r   r5   r6   r   r   r4   
test_casestest_cases_max_diffr2   r3   r{   r   rC   torch_outputsmax_diffmax_diff_allir   r   r   verify_onnx   sL   		 $ $
,0
z&T5EncoderDecoderInitHelper.verify_onnxN)TTFF)r   )r'   r(   r)   rE   r   r+   r5   rc   rF   r   r1   r   r   r.   r   r   r   r   r   rG   `   sJ     rG   )*loggingri   sysrg   pathlibr   typingr   r   r   r   ro   r+   past_helperr   
t5_decoderr   r   r   r	   transformersr
   r   onnxruntimer   rj   rB   rk   dirname__file__
onnx_modelr   torch_onnx_export_helperr   	getLoggerr'   r   r,   r-   r   r1   rG   r   r   r   r   <module>   s*    
$