o
    ;ήcNB                     @   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 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 ejejejedd d d	lmZ 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ej#j$Z&G dd dZ'G dd dZ(dS )    N)Path)ListUnion)PastKeyValuesHelper)T5EncoderInputs)	MT5ConfigT5Config)InferenceSessionz..)
TypeHelper)	OnnxModel)torch_onnx_exportc                	       sb   e Zd ZdZ	ddejjdejjdeee	f de
f fddZd	ejd
ejdejfddZ  ZS )T5DecoderInitz~A T5 decoder with LM head to create initial past key values.
    This model is only called once during starting decoding.
    Ndecoderlm_headconfigdecoder_start_token_idc                    s<   t    || _|| _|| _|d ur|| _d S | jj| _d S N)super__init__r   r   r   r   )selfr   r   r   r   	__class__ T/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/t5/t5_decoder.pyr   $   s   

zT5DecoderInit.__init__decoder_input_idsencoder_attention_maskencoder_hidden_statesc                 C   s   |d u r|j d }tj|dftj|jd| j }| j|||ddd}|j}|j}|| j	j
d  }| |}t|\}	}
||	|
fS )Nr      dtypedeviceT)	input_idsr   r   	use_cachereturn_dict      )shapetorchoneslongr    r   r   last_hidden_statepast_key_valuesr   d_modelr   r   group_by_self_or_cross)r   r   r   r   
batch_sizedecoder_outputssequence_outputpresent_key_values	lm_logits	past_self
past_crossr   r   r   forward3   s.   
	

zT5DecoderInit.forwardr   )__name__
__module____qualname____doc__r&   nnModuler   r   r   intr   TensorFloatTensorr4   __classcell__r   r   r   r   r      s&    	
r   c                       s(   e Zd ZdZ fddZdd Z  ZS )	T5Decoderz-A T5 decoder with LM head and past key valuesc                    s    t    || _|| _|| _d S r   )r   r   r   r   r   )r   r   r   r   r   r   r   r   Y   s   

zT5Decoder.__init__c                 G   sb   t || jj}| j||||ddd}|j}|j}|| jjd  }| |}	t 	|\}
}|	|
fS )NT)r!   r*   r   r   r"   r#   r$   )
r   group_by_layerr   
num_layersr   r)   r*   r+   r   r,   )r   r   r   r   pastr*   r.   r/   r0   r1   present_self_r   r   r   r4   _   s   	
zT5Decoder.forward)r5   r6   r7   r8   r   r4   r>   r   r   r   r   r?   V   s    r?   c                   @   sh   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defddZdefddZdd ZdS )T5DecoderInputsNc                 C   s   || _ || _|| _|| _d S r   )r   r   r   r*   )r   r   r   r   r*   r   r   r   r   y   s   
zT5DecoderInputs.__init__Fr   r-   encode_sequence_lengthpast_decode_sequence_lengthr    float16use_int32_inputsc                 C   s  | j }| j}| j}	| j}
| j}d}tjd|
d ||f|rtjntj|d}t	j
|||
||d}|r4tjntj}tj|||||d}|dkr|||||g}||||g}g }td|	 D ]}|tj|||d qYtd|	 D ]}|tj|||d qmnd}t||j||S )aZ  Create dummy inputs for T5Decoder.

        Args:
            decoder: decoder
            batch_size (int): batch size
            encode_sequence_length (int): sequence length of input_ids for encoder
            past_decode_sequence_length (int): past sequence length of input_ids for decoder
            device (torch.device): device of output tensors
            float16 (bool): whether the model uses float32 or float16 in input
            use_int32_inputs(bool): whether use int32 instead of int64 for some inputs

        Returns:
            T5DecoderInputs: dummy inputs for decoder
        r   r   )lowhighsizer   r    )rI   r      N)r+   	num_headsrA   
vocab_sized_kvr&   randintint32int64r   create_dummyrH   float32randrangeappendrE   attention_mask)r   r-   rF   rG   r    rH   rI   hidden_sizenum_attention_headsrA   rO   	head_sizesequence_lengthr   encoder_inputs
float_typeencoder_hidden_stateself_attention_past_shapecross_attention_past_shaperB   rD   r   r   r   rT      s^   zT5DecoderInputs.create_dummyreturnc                 C   s&   | j | j| jg}| jr|| j |S r   )r   r   r   r*   extend)r   
input_listr   r   r   to_list   s   zT5DecoderInputs.to_listc                 C   sD   | j jtjd}| jrdd | jD nd }t| j | j ||S )Nr   c                 S   s   g | ]	}|j tjd qS )rg   )tor&   rU   ).0pr   r   r   
<listcomp>   s    z+T5DecoderInputs.to_fp32.<locals>.<listcomp>)	r   rh   r&   rU   r*   rE   r   cloner   )r   r`   rB   r   r   r   to_fp32   s   zT5DecoderInputs.to_fp32r   )FF)r5   r6   r7   r   staticmethodr   r   r   r;   r&   r    boolrT   r   rf   rm   r   r   r   r   rE   x   s.    

S
rE   c                   @   s   e Zd Ze			ddeeef 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eef dedejde
def
ddZdS )T5DecoderHelperTFr   r    onnx_model_pathverboseuse_external_data_formatrI   c                 C   s  t | ttfs	J tj| jddt | trdnd||d}| }tj| jj	dd}tj| jj	dd}	|	d	d| jj	  }
t | trC|ng }t | trL|
n|	}d
g| }dg}|
d |
d || ddiddddddddid}|D ]}dd|v rdndd||< qx|D ]!}d|v rddd||< qt | trddd||< qddi||< qt|jjddd t I}tj|d}t|jjddd t| t||r|n|d|||dd||d |rtj|dd}tj||ddd W d	   d	S W d	   d	S 1 sw   Y  d	S )a  Export decoder to ONNX

        Args:
            decoder (Union[T5Decoder, T5DecoderNoPastState]): decoder object
            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.
            use_int32_inputs (bool, optional): use int32 inputs
        rM         r   )r-   rF   rG   r    rI   F)presentTNlogitsr!   r   r   r-   rF   )r   r   )r!   r   r   rw   r   rG   )r   rM   crosszpast_decode_sequence_length + 1)parentsexist_okzdecoder.onnx   )
argsfexport_paramsinput_namesoutput_namesdynamic_axesopset_versiondo_constant_foldingrs   rr   )load_external_data)save_as_external_dataall_tensors_to_one_file)
isinstancer?   r   rE   rT   r   rf   r   get_past_namesrA   rX   rd   r   parentmkdirtempfileTemporaryDirectoryospathjoinr   tupleonnx
load_modelr   save)r   r    rq   rr   rs   rI   inputsre   
past_namespresent_namespresent_self_namesinput_past_namesoutput_present_namesr   r   r   nametmp_dir_nametemp_onnx_model_pathmodelr   r   r   export_onnx   s   







$zT5DecoderHelper.export_onnxr   c                 C   s   t d t|j  t|j  t|j  d}|jrVt	|jd dks1J t
t	|jd }t|}t|jD ]\}}t|  ||| < qD| d|}|S )zRun inference of ONNX model.zstart onnxruntime_inference)r!   r   r      r   N)loggerdebugnumpyascontiguousarrayr   cpur   r   r*   lenr;   r   r   	enumeraterun)ort_sessionr   
ort_inputsrA   r   ipast_tensorort_outputsr   r   r   onnxruntime_inferencee  s   

z%T5DecoderHelper.onnxruntime_inferencer   r   r   	max_casesc                 C   s  t |ddk}g d}g }|d| D ]\}}	}
t| tr d}
tj| j||	|
|||d}|  }t	
  | | }W d   n1 sFw   Y  t||}tt|d   |d  }|}td|  td| jj D ](}tt|d	 |   |d	|   }td
| d|  t||}qut| trtd| jj D ].}tt|d |   |d	d| jj  |   }td| d|  t||}q|| td| d|	 dd|
 d|   q|S )zQCompare the result from PyTorch and OnnxRuntime to verify the ONNX model is good.r   ztensor(float16)))r      rt   )r   rM   ru   )rt   r   r   )   ru   rM   Nr   )r    rH   rI   zlogits max_diff=rM   r   zself attention past state z
 max_diff=zcross attention past state zbatch_size=z, encode_sequence_length=z, zpast_decode_sequence_length=z, max_diff=)r
   get_input_typer   r   rE   rT   r   rm   rf   r&   no_gradrp   r   r   amaxabsr   r   r   rW   rA   maxrX   info)r   r   r    rI   r   rH   
test_casestest_cases_max_diffr-   rF   rG   r   re   torch_outputsr   max_diffmax_diff_allr   r   r   r   verify_onnxz  sZ   	



$,
0
zT5DecoderHelper.verify_onnxN)TFF)r   )r5   r6   r7   rn   r   r?   r   r&   r    strro   r   rE   r   r	   r;   r   r   r   r   r   rp      sB    
u
rp   ))loggingr   sysr   pathlibr   typingr   r   r   r   r&   past_helperr   
t5_encoderr   transformersr   r   onnxruntimer	   r   rX   r   dirname__file__io_binding_helperr
   
onnx_modelr   torch_onnx_export_helperr   	getLoggerr5   r   r9   r:   r   r?   rE   rp   r   r   r   r   <module>   s,    
7"v