o
    ;ήc(                     @   s   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
mZmZmZ d dlmZ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! d dl"m#Z# e $e%Z&g dZ'g dZ(G dd dZ)dS )    N)Path)DictListUnion)	T5DecoderT5DecoderHelperT5DecoderInit)	T5EncoderT5EncoderHelper)T5EncoderDecoderInitT5EncoderDecoderInitHelper)MT5ForConditionalGenerationT5ForConditionalGeneration)InferenceSessionz..)float_to_float16_max_diff)	OnnxModel)optimize_model)zt5-smallzt5-basezt5-largezt5-3bzt5-11b)zgoogle/mt5-smallzgoogle/mt5-basezgoogle/mt5-largezgoogle/mt5-xlzgoogle/mt5-xxlc                   @   s.  e Zd Ze		d*dededededef
dd	Ze	
	d+dededejdedede	eej
jf fddZe	
		
	d,deeeeef dejdededededefddZeg dfdedee fddZe		
d-deded ed!ed"eded#efd$d%Zedeeeeef d&edejdefd'd(Zd)S ).T5Helper F
output_dirmodel_name_or_pathsuffix
new_folderreturnc                 C   s^   |}t j|rt|jd }n|dd  ||7 }|r$t j| |n| }t j||d S )a  Build onnx path

        Args:
            output_dir (str): output directory
            model_name_or_path (str): pretrained model name, or path to the model checkpoint
            suffix (str, optional): suffix like "_encoder" or "_decoder_fp16" will be appended to file name. Defaults to None.
            new_folder (bool, optional): create a new directory for the model. Defaults to False.

        Returns:
            str: path of onnx model
        /z.onnx)ospathisdirr   partssplitjoin)r   r   r   r   
model_name	directory r$   S/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/t5/t5_helper.pyget_onnx_path!   s   zT5Helper.get_onnx_pathTt5	cache_dirdevicemerge_encoder_and_decoder_init
model_typec           
      C   s   |dkrt j| |d}n|dkrtj| |d}ntdt|j|j|j}| 	| |r@t
|j|j|j|jdd}||dS t|j|j}| 	| t|j|j|j}	|	 	| |||	dS )	ab  Load model given a pretrained name or path, then build models for ONNX conversion.

        Args:
            model_name_or_path (str): pretrained model name or path
            cache_dir (str): cache directory
            device (torch.device): device to run the model
            merge_encoder_and_decoder_init (bool, optional): Whether merge encoder and decoder initialization into one ONNX model. Defaults to True.
            is_mt5 (bool, optional): whether the model is MT5 instead of T5
        Returns:
            Dict[str, torch.nn.Module]: mapping from name to modules for ONNX conversion.
        r'   )r(   mt5z only support mode_type=t5 or mt5N)decoder_start_token_id)encoder_decoder_initdecoder)encoderr/   decoder_init)r   from_pretrainedr   
ValueErrorr   r/   lm_headconfigevaltor   r0   r	   r   )
r   r(   r)   r*   r+   modelr/   r.   r0   r1   r$   r$   r%   
load_model>   s0   
zT5Helper.load_modelr8   onnx_model_pathverboseuse_external_data_formatuse_decoder_input_idsuse_int32_inputsc              	   C   s^   t | trt| ||||| d S t | tr#t| |||||| d S t| ||||| d S )N)
isinstancer	   r
   export_onnxr   r   r   )r8   r)   r:   r;   r<   r=   r>   r$   r$   r%   r@   o   s6   



zT5Helper.export_onnx)Pow
ReduceMeanAddSqrtDivMulSoftmaxRelu
onnx_modelop_block_listc                 C   sT  t dd |  D }t |}||}td| d|  |  jd j}d}|  }||v s3J || }d}	|j	dkrq|}	td	|j  d}
|j
D ]}| |}
|
dur[ nqNt|
}td
|j d|  |dk }ntd|j	 d|j  g }g }|s|	dur|g}|	jg}||||d}td|  | jdddi| |S )a  Convert model to mixed precision.
           It detects whether original model has fp16 precision weights, and set parameters for float16 conversion automatically.
        Args:
            onnx_model (OnnxModel): optimized ONNX model
            op_block_list (List[str], optional): . Defaults to ["Pow", "ReduceMean", "Add", "Sqrt", "Div", "Mul", "Softmax", "Relu"]
        Returns:
            parameters(dict): a dictionary of parameters used in float16 conversion
        c                 S   s   g | ]}|j qS r$   )op_type).0noder$   r$   r%   
<listcomp>   s    z1T5Helper.auto_mixed_precision.<locals>.<listcomp>z	fp32 op: z
 fp16 op: r   FNMatMulz#Found last MatMul node for logits: z3max diff of converting weights in last MatMul node z: gư>z-Failed to find MatMul node for logits. Found z	 of node )keep_io_typesrJ   node_block_listforce_fp16_initializersz!auto_mixed_precision parameters: use_symbolic_shape_inferTr$   )setnodes
differenceloggerinfographoutputnameoutput_name_to_noderK   inputget_initializerr   debugwarningconvert_float_to_float16)rI   rJ   op_full_setfp32_op_setfp16_op_setlogits_output_nameis_weight_fp16_precisionr\   rM   last_matmul_nodeinitializerr]   max_diffrP   rQ   
parametersr$   r$   r%   auto_mixed_precision   sH   




zT5Helper.auto_mixed_precisionoptimized_model_path
is_float16num_attention_headshidden_sizerk   c              	   C   sJ   t | d||dddd}|r|rt| n|jdd |j||dd dS )	zHOptimize ONNX model with an option to convert it to use mixed precision.bertr   NF)r+   	num_headsro   	opt_leveloptimization_optionsuse_gpu)cast_input_outputT)all_tensors_to_one_file)r   r   rk    convert_model_float32_to_float16save_model_to_file)r:   rl   rm   rn   ro   r<   rk   mr$   r$   r%   optimize_onnx   s   	zT5Helper.optimize_onnxort_sessionc                 C   sD   t | trt| |||S t | trt| |||S t| |||S )zQCompare the result from PyTorch and OnnxRuntime to verify the ONNX model is good.)r?   r	   r
   verify_onnxr   r   r   )r8   r{   r)   r>   r$   r$   r%   r|      s
   

zT5Helper.verify_onnxN)r   F)Tr'   )TFTF)FT)__name__
__module____qualname__staticmethodstrboolr&   torchr)   r   nnModuler9   r   r	   r   r   r   r@   r   r   rk   intrz   r   r|   r$   r$   r$   r%   r       s    0&Gr   )*loggingr   syspathlibr   typingr   r   r   r   
t5_decoderr   r   r   
t5_encoderr	   r
   t5_encoder_decoder_initr   r   transformersr   r   onnxruntimer   r   appendr!   dirname__file__float16r   rI   r   	optimizerr   	getLoggerr}   rW   PRETRAINED_T5_MODELSPRETRAINED_MT5_MODELSr   r$   r$   r$   r%   <module>   s&    
