o
    ;ήcY                     @   s   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 d dlmZm	Z	 d dl
Z
d dlZd dl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Union)	MT5ConfigT5Config)InferenceSessionz..)	OnnxModel)torch_onnx_exportc                       s6   e Zd ZdZdeeef f fddZdd Z  Z	S )	T5Encoderz-T5 encoder outputs only the last hidden stateconfigc                    s   t    || _|| _d S N)super__init__encoderr   )selfr   r   	__class__ T/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/t5/t5_encoder.pyr       s   

zT5Encoder.__init__c                 C   s   |  ||d S )Nr   )r   r   	input_idsattention_maskr   r   r   forward%   s   zT5Encoder.forward)
__name__
__module____qualname____doc__r   r   r   r   r   __classcell__r   r   r   r   r
      s    r
   c                   @   sJ   e Zd Zdd Ze	ddedededejdef
d	d
Z	de
fddZdS )T5EncoderInputsc                 C   s   || _ || _d S r   r   r   r   r   r   r   r   *   s   
zT5EncoderInputs.__init__F
batch_sizesequence_length
vocab_sizedeviceuse_int32_inputsc           
      C   s   |rt jnt j}t jd|d | |f||d}t j| |g||d}|dkr;t| D ]}td|d }	d||d|	f< q(t||S )aI  Create dummy inputs for T5 encoder.

        Args:
            batch_size (int): batch size
            sequence_length (int): sequence length
            vocab_size (int): vocabulary size
            device (torch.device): device of output tensors

        Returns:
            T5EncoderInputs: dummy inputs for encoder
        r      )lowhighsizedtyper#   )r)   r#      N)torchint32int64randintonesrangerandomr   )
r    r!   r"   r#   r$   r)   r   r   ipadding_positionr   r   r   create_dummy.   s   
zT5EncoderInputs.create_dummyreturnc                 C   s   dd | j | jfD }|S )Nc                 S   s   g | ]}|d ur|qS r   r   ).0vr   r   r   
<listcomp>O   s    z+T5EncoderInputs.to_list.<locals>.<listcomp>r   )r   
input_listr   r   r   to_listN   s   zT5EncoderInputs.to_listNF)r   r   r   r   staticmethodintr+   r#   boolr4   r   r:   r   r   r   r   r   )   s     r   c                   @   sr   e Zd Ze			ddedej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fddZdS )T5EncoderHelperTFr   r#   onnx_model_pathverboseuse_external_data_formatr$   c                 C   s  | j }tjdd|j||d}t|jjddd t [}t	j
|d}	t|	jjddd t| t| |r9|	n|dddgd	gd
ddd
ddd
ddddd||d |rotj|	dd}
tj|
|ddd W d   dS W d   dS 1 szw   Y  dS )a  Export encoder to ONNX

        Args:
            encoder (T5Encoder): encoder object
            device (torch.device): device of encoder 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.
        r*      r    r!   r"   r#   r$   T)parentsexist_okzencoder.onnxr   r   hidden_statesr    r!   )r   r%   )r   r   rG      )
argsfexport_paramsinput_namesoutput_namesdynamic_axesopset_versiondo_constant_foldingrB   rA   )load_external_data)save_as_external_dataall_tensors_to_one_fileN)r   r   r4   r"   r   parentmkdirtempfileTemporaryDirectoryospathjoinr	   tupler:   onnx
load_modelr   save)r   r#   r@   rA   rB   r$   r   encoder_inputstmp_dir_nametemp_onnx_model_pathmodelr   r   r   export_onnxT   sN   


"zT5EncoderHelper.export_onnxinputsc                 C   s6   t |j   t |j   d}| d|S )zRun inference of ONNX model.r   N)numpyascontiguousarrayr   cpur   run)ort_sessionrd   
ort_inputsr   r   r   onnxruntime_inference   s   z%T5EncoderHelper.onnxruntime_inferencerb   ri   c           	      C   sh   t jdd| jj||d}| }| | }t||}tt	|
  |d  }td|  |S )zQCompare the result from PyTorch and OnnxRuntime to verify the ONNX model is good.rC      rD   r   z	max_diff=)r   r4   r   r"   r:   r?   rk   re   amaxabsrg   loggerinfo)	rb   ri   r#   r$   rd   r9   torch_outputsort_outputsmax_diffr   r   r   verify_onnx   s    zT5EncoderHelper.verify_onnxN)TFFr;   )r   r   r   r<   r
   r+   r#   strr>   rc   r   rk   r   rt   r   r   r   r   r?   S   s>    :	r?   )#loggingrX   r1   sysrV   pathlibr   typingr   r   re   r\   r+   transformersr   r   onnxruntimer   rY   appendrZ   dirname__file__
onnx_modelr   torch_onnx_export_helperr	   	getLoggerr   ro   nnModuler
   r   r?   r   r   r   r   <module>   s&    
*