o
    ;ήc[                     @   sP  d 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 ddlm	Z	 ddl
mZmZ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mZmZ ddlmZmZmZmZmZmZm Z m!Z! dd	l"m#Z#m$Z$m%Z%m&Z& ej'(ej')ej'*e+d
d ddl,m-Z- ddl.m/Z0 ej'(ej')ej'*e+d
d ddlm1Z1 ddl2m3Z4 ddl5m6Z6m7Z7 ddl8m9Z9 e:dZ;G dd deZ<dSdeee=  dej>fddZ?dej>fddZ@dej>fddZAdTd e=d!eBfd"d#ZCd$e=d%eBde$fd&d'ZDd(ejd)efd*d+ZEd(ejd)efd,d-ZFd(ejd)efd.d/ZG	0	1dUd2ed3ed4e=d5eHfd6d7ZId8ed9efd:d;ZJ	1dVd(ed5eHdee fd<d=ZKe<jLfdej>d>e<fd?d@ZMdej>dAeee f dBejNdCejNdDeHdEeHdFeeeH  dee=ef fdGdHZOdWdej>dJeee=  dKeBfdLdMZPdSdej>dJeee=  fdNdOZQdXdeee=  dJeee=  fdPdQZ/eRdRkre/  dS dS )Ya  
This converts GPT2 or T5 model to onnx with beam search operator.

Example 1: convert gpt2 model with beam search:
    python convert_generation.py -m gpt2 --output gpt2_beam_search.onnx

Example 2: convert T5 model with beam search in two steps:
    cd ./models/t5
    python convert_to_onnx.py -m t5-small
    cd ../..
    python convert_generation.py -m t5-small --model_type t5                                           --decoder_onnx ./models/t5/onnx_models/t5-small_decoder.onnx                                    --encoder_decoder_init_onnx ./models/t5/onnx_models/t5-small_encoder_decoder_init.onnx          --output ./models/t5/onnx_models/t5_small_beam_search.onnx

Example 3: convert T5 model with beam search. All in one step:
    python convert_generation.py -m t5-small --model_type t5 --output ./models/t5/onnx_models/t5_small_beam_search.onnx

Example 4: convert MT5 model with external data file like mt5-base-beamsearch.onnx.data in below example.
    python convert_generation.py -m google/mt5-base --model_type mt5 --output mt5-base-beamsearch.onnx -e

Example 5: convert gpt2 model with greedy search:
    python convert_generation.py -m gpt2 --output gpt2_greedy_search.onnx --num_beams 1 --num_return_sequences 1
    N)Enum)Path)AnyDictListOptionalUnion)	Precision)
GraphProto
ModelProtoTensorProto)
GPT2ConfigGPT2LMHeadModelGPT2Tokenizer	MT5ConfigMT5ForConditionalGenerationT5ConfigT5ForConditionalGenerationT5Tokenizer)GraphOptimizationLevelInferenceSessionSessionOptionsget_available_providersmodelsgpt2)PRETRAINED_GPT2_MODELS)maint5)setup_logger)export_onnx_models)PRETRAINED_MT5_MODELSPRETRAINED_T5_MODELS)	OnnxModel c                   @   s   e Zd ZdZdZdd ZdS )GenerationTypebeam_searchgreedy_searchc                 C   s   | j S N)value)self r*   R/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/convert_generation.py__str__J   s   zGenerationType.__str__N)__name__
__module____qualname__
BEAMSEARCHGREEDYSEARCHr,   r*   r*   r*   r+   r$   F   s    r$   argvreturnc                 C   sl  t  }|d}|jdddtddtt t  d |jdd	td
g dddg d d |jdd	tt	j
dddd |jdd	tddd |jdd	tddd |jdd	ddd |jd	d |d}|jddtdd |jd d!d	ttjtjtjg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	dd-d |jd	d. |d/}|jd0d	dd1d |jd	d2 |jd3d	dd4d |jd	d5 |jd6d	dd7 |jd	d8 |jd9td	d:d;d< |jd=d	dd>d |jd	d? |jd@d	ddAd |jd	dB |jdCd	ddDd |jd	dE |dF}|jdGtd	dHdId< |jdJtd	dKdLd< |jdMtd	dNdOd< |jdPtd	dHdQd< |jdRtd	dHdSd< |jdTtd	dHdUd< |jdVtd	dWdXd< |dY}|jdZd	dd[d |jd	d\ |jd]d	dd^d |jd	d_ |jd`d	ddad |jd	db |jdcd	tdHddd |jded	ddfd |jd	dg || }|S )hzParse arguments

    Args:
        argv (Optional[List[str]], optional): _description_. Defaults to None.

    Returns:
        argparse.Namespace: Parsed arguments.
    zInput optionsz-m--model_name_or_pathTzEPytorch model checkpoint path, or pretrained model name in the list: , )requiredtypehelpz--model_typeFr   )r   r   mt5z*Model type (default is gpt2) in the list: )r6   r7   defaultchoicesr8   z--cache_dir.cache_modelsz%Directory to cache pre-trained models)r6   r7   r:   r8   z--decoder_onnxr#   zLPath of onnx model for decoder. Specify it when you have exported the model.z--encoder_decoder_init_onnxzgPath of ONNX model for encoder and decoder initialization. Specify it when you have exported the model.z	--verbose
store_truezPrint more information)r6   actionr8   )verbosezOutput options--outputz,Output path for onnx model with beam search.z-p--precisionzTPrecision of model to run. fp32 for full precision, fp16 for half or mixed precisionz-e--use_external_data_formatz!save external data for model > 2G)use_external_data_formatz-sz--run_shape_inferencezrun shape inference)run_shape_inferencez-iz--disable_shared_initializerszado not share initializers in encoder and decoder. It will increase memory usage of t5/mt5 models.)disable_shared_initializersz6Beam search parameters that stored in the output modelz--output_sequences_scoreszoutput sequences scores)output_sequences_scoresz--output_token_scoreszoutput token scores)output_token_scoresz--early_stopping)r6   r?   )early_stoppingz--no_repeat_ngram_sizer   zNo repeat ngram size)r7   r6   r:   r8   z--vocab_maskz\Enable vocab_mask. This mask applies only to every generated token to filter some bad words.)
vocab_maskz--prefix_vocab_maskzeEnable prefix_vocab_mask. This mask can be used to filter bad words in the first generated token only)prefix_vocab_maskz--custom_attention_maskz]Enable custom_attention_mask. This mask can be used to replace default encoder attention mask)custom_attention_maskzYBeam search parameters not stored in the output model, for testing parity and performancez--min_length   zMin sequence lengthz--max_length2   zMax sequence lengthz--num_beams   z	Beam sizez--num_return_sequencesz&Number of return sequence <= num_beamsz--length_penaltyz<Positive. >1 to penalize and <1 to encourage short sentence.z--repetition_penaltyz-Positive. >1 to penalize and <1 to encourage.z--vocab_sizezIVocab_size of the underlying model used to decide the shape of vocab maskz0Other options for testing parity and performance	--use_gpuz)use GPU for inference. Required for fp16.)use_gpuz--disable_parityzdo not run parity test)disable_parityz--torch_performanceztest PyTorch performance)torch_performancez--total_runsz4Number of times of inference for latency measurementz--save_test_dataz.save test data for onnxruntimer_perf_test tool)save_test_data)argparseArgumentParseradd_argument_groupadd_argumentstrjoinr   r!   r    ospathset_defaultsr	   FLOAT32FLOAT16intfloat
parse_args)r2   parserinput_groupoutput_groupmodel_groupbeam_parameters_group
test_groupargsr*   r*   r+   parse_argumentsN   s  	
		






rk   rj   c                 C   s   | j }d|d| jdd| jtjkrdndddd	d
ddg}| jr#|d | jr+|d | jtjkr?| js8J d|	g d | j
rJtd|  t|d dS )zqConvert GPT-2 model to onnx

    Args:
        args (argparse.Namespace): arguments parsed from command line
    r4   rA   z--optimize_onnxrB   fp32fp16z--test_runs1z--test_cases10z--use_int32_inputsz--overwriterQ   rC   zEfp16 or mixed precision model cannot run in CPU. Please add --use_gpu)z--op_block_listAddLayerNormalizationFastGeluzarguments for convert_to_onnx:)r2   N)model_name_or_pathdecoder_onnx	precisionr	   r_   rR   appendrD   r`   extendr@   loggerinfoconvert_gpt2_to_onnx)rj   
model_name	argumentsr*   r*   r+   gpt2_to_onnxB  s2   

r}   c                 C   sx   t | j| jt| jj| j| jd| jdddddd| j	d}t
d|d   t
d|d   |d | _|d | _dS )	znConvert T5 model to onnx

    Args:
        args (argparse.Namespace): arguments parsed from command line
    FT)rR   rD   optimize_onnxru   r@   use_decoder_start_tokenmerge_encoder_and_decoder_init	overwritedisable_auto_mixed_precisionuse_int32_inputs
model_typezonnx model for encoder: r   zonnx model for decoder: rM   N)export_t5_onnx_modelsrs   	cache_dirr   outputparentrR   rD   ru   r   rx   debugencoder_decoder_init_onnxrt   )rj   pathsr*   r*   r+   
t5_to_onnxk  s(   

r   T	onnx_pathrD   c                 C   sP   ddl m} tj| dd}|j|ddd}|r!tj|| |d d	S td d	S )
zShape inference on an onnx file, which will be overwritten.

    Args:
        onnx_path (str): Path of onnx model
        use_external_data_format(bool): output tensors to external data or not.
    r   )SymbolicShapeInferenceTload_external_dataF)
auto_mergeguess_output_rank)save_as_external_dataz4Failed to run symbolic shape inference on the model.N)	&onnxruntime.tools.symbolic_shape_inferr   onnx
load_modelinfer_shapesr"   saverx   warning)r   rD   r   modeloutr*   r*   r+   shape_inference  s   r   
model_pathrR   c                 C   sR   t  }tj|_|rddgndg}|r dt vrtdtd t| ||d}|S )a,  Create OnnxRuntime session.

    Args:
        model_path (str): onnx model path
        use_gpu (bool): use GPU or not

    Raises:
        RuntimeError: CUDAExecutionProvider is not available when --use_gpu is specified.

    Returns:
        onnxruntime.InferenceSession: The created session.
    CUDAExecutionProviderCPUExecutionProviderz5CUDAExecutionProvider is not available for --use_gpu!zuse CUDAExecutionProvider)	providers)	r   r   ORT_DISABLE_ALLgraph_optimization_levelr   RuntimeErrorrx   ry   r   )r   rR   sess_optionsexecution_providersort_sessionr*   r*   r+   create_ort_session  s   

r   graphru   c              	   C   s  t j|k}t| j}|d }|dksJ g ddd t|D  }t| jt|kr9tdt| dt| j t|D ]E\}}| j| j|krZtd| d	| d| j| j tj	}|dkri|rftjntj
}| j| jjj}	|	|krtd| d
| d|	 q=td dgdd t|D  }
t| jt|
krtdt|
 dt| j t|
D ]>\}}| j| j|krtd| d	| d| j| j |rtjntj
}| j| jjj}||krtd| d
| d| qtd dS )a  Verify GPT-2 subgraph

    Args:
        graph (onnx.GraphProto): onnx graph of GPT-2
        precision (Precision): Precision (FLOAT16 or FLOAT32) of the model.

    Raises:
        ValueError: Number of inputs not expected.
        ValueError: Input name is not expected.
        ValueError: Input data type is not expected.
        ValueError: Number of outputs not expected.
        ValueError: Output name is not expected.
        ValueError: Output data type is not expected.
       rM   )	input_idsposition_idsattention_maskc                 S      g | ]}d | qS )past_r*   .0ir*   r*   r+   
<listcomp>      z(verify_gpt2_subgraph.<locals>.<listcomp> Number of inputs expected to be . Got Input  is expected to be $ is expected to have onnx data type z:Verifying GPT-2 graph inputs: name and data type are good.logitsc                 S   r   )present_r*   r   r*   r*   r+   r     r   !Number of outputs expected to be Output z;Verifying GPT-2 graph outputs: name and data type are good.N)r	   r`   leninputrange
ValueError	enumeratenamer   INT32FLOATr7   tensor_type	elem_typerx   ry   r   )r   ru   
is_float16input_countlayer_countexpected_inputsr   expected_inputexpected_type
input_typeexpected_outputsexpected_outputoutput_typer*   r*   r+   verify_gpt2_subgraph  s>   

"
"
r   c              	   C   s8  t j|k}|r
tjntj}t| j}|d d }|dksJ g d}t|D ]}|d|  |d|  q&t|D ]}|d|  |d|  q=t| jt|krhtd	t| d
t| j t	|D ]?\}}| j| j
|krtd| d| d
| j| j
 |dk rtjn|}	| j| jjj}
|
|	krtd| d|	 d
|
 qldg}t|D ]}|d|  |d|  qt| jt|krtdt| d
t| j t	|D ]7\}}| j| j
|krtd| d| d
| j| j
 | j| jjj}||krtd| d| d
| qdS )  Verify T5 decoder subgraph

    Args:
        graph (onnx.GraphProto): onnx graph of T5 decoder
        precision (Precision): Precision (FLOAT16 or FLOAT32) of the model.

    Raises:
        ValueError: Number of inputs not expected.
        ValueError: Input name is not expected.
        ValueError: Input data type is not expected.
        ValueError: Number of outputs not expected.
        ValueError: Output name is not expected.
        ValueError: Output data type is not expected.
    r   rO   rM   )r   encoder_attention_maskencoder_hidden_statespast_key_self_past_value_self_past_key_cross_past_value_cross_r   r   r   r      r   r   present_key_self_present_value_self_r   r   N)r	   r`   r   r   r   r   r   rv   r   r   r   r   r7   r   r   r   )r   ru   r   
float_typer   r   r   r   r   r   r   r   r   r   r*   r*   r+   verify_t5_decoder_subgraph  sH   

""
r   c              	   C   s  t j|k}t| jd d }|dksJ g d}t| jt|kr0tdt| dt| j t|D ]9\}}| j| j|krQtd| d| d| j| j tj	}| j| j
jj}||krmtd| d	| d| q4d
dg}	t|D ]}|	d|  |	d|  qvt|D ]}|	d|  |	d|  qt| jt|	krtdt|	 dt| j t|	D ]>\}}
| j| j|
krtd| d|
 d| j| j |rtjntj}| j| j
jj}||krtd| d	| d| qtd dS )r   r   rO   rM   )encoder_input_idsr   decoder_input_idsr   r   r   r   r   r   r   r   r   present_key_cross_present_value_cross_r   r   zMT5 encoder graph verified: name and data type of inputs and outputs are good.N)r	   r`   r   r   r   r   r   r   r   r   r7   r   r   r   rv   r   rx   ry   )r   ru   r   r   r   r   r   r   r   r   r   r   r*   r*   r+   'verify_t5_encoder_decoder_init_subgraph9  s@   
""r   shared_   graph1graph2shared_prefixmin_elementsc                 C   s  i }i }g }g }g }| j D ]L}	|	jrt|	j|ksq|j D ];}
|
jr)t|
j|ks*qt|	|
rX||
j ||	j< ||	 |
j|vrV||
j }|||
j< ||
 ||  nqqtd|  | j	D ]}t
t|jD ]}|j| |v rtd|j|  qnqe|j	D ]}t
t|jD ]}|j| |v rtd|j|  qq|D ]}|j | q|jD ]}|j|v r||j |_q|j	D ]4}t
t|jD ]*}|j| |v r||j|  }td|j d| d|j|  d|  ||j|< qq|D ]}| j | q| jD ]}|j|v r||j |_q| j	D ]7}t
t|jD ],}|j| |v rM||j|  }td|j d| d|j|  d|  ||j|< q"q|D ]	}||j |_qS|D ] }tj|j}tj|j|j|}| j| |j| q_|S )	a  Remove initializers with same value from two graphs.

    Args:
        graph1 (GraphProto): the first graph to process
        graph2 (GraphProto): the second graph to process
        shared_prefix (str): add prefix to the shared initializers among two graphs
        min_elements (int, optional): minimal number of elements for initializers to be considered. Defaults to 1024.
    zshared initializers:zname is found in graph 1: zname is found in graph 2: zgraph 2 rename node z input z from z to zgraph 1 rename node )initializerdimssumr"   has_same_valuer   rv   rx   r   noder   r   r   r   remove
value_infor   numpy_helperto_arrayshapehelpermake_tensor_value_info	data_type)r   r   r   r   mapping_initializers_1mapping_initializers_2shared_initializers_1shared_initializers_2shared_initializers_namesinitializer1initializer2shared_namer   jr   r   new_namer   r*   r*   r+   remove_shared_initializers}  s   












*


*
r   encoder_modeldecoder_modelc                 C   sL   t | }t |}|d |d |  |  t|jj|jjd}|S )Ne_d_s_)r"   add_prefix_to_namesremove_duplicated_initializerr   r   r   )r  r  encoderdecoderinitializersr*   r*   r+   get_shared_initializers  s   

r  c                 C   s   g }| j D ]}|jrt|j|ksq|| q|D ]}| j | q|D ]}tj|j}tj	
|j|j|}| j| q%|S )a^  Remove initializers of a graph, when they have number of elements larger than a threshold.

    Args:
        graph (GraphProto): the graph.
        min_elements (int, optional): minimal number of elements for initializers to be considered. Defaults to 1024.

    Returns:
        List[TensorProto]: initializers that are removed from the graph.
    )r   r   r   rv   r   r   r   r   r   r   r   r   r   r   )r   r   moved_initializerstensorr   r   r   r*   r*   r+   move_initializers  s   
r  generation_typec           "   	   C   s<  | j dk}|tjk}|r |std| jrtd| jr td|re| jr6tj	| jr6t
d| j  nQ| jsRd| jtjkrCdnd}tt| jj| | _t
d	| j d
| j d t|  n"| jry| jryt
d| j d| j  nt
d| j d t|  | jrt
d| j d t| j| j |rtj| j| jd}n| j dkrtj| j| jd}n	tj| j| jd}| j rt
d|  |j!}|r|j!n|j"}|j#}| j#dkr| j#}t$j%| jdd}	| j  d|	j&_'| j dkrt(|	j&| j nt)|	j&| j |sg dng d}
| j*r|
+d n|
+d | j,r(|
+d n|
+d | j-r7|
+d n|s?|
+d dg}| jrK|+d  | jr\| jsWJ d!|+d" |smt$j.j/d#|
|d$| j  d%nt$j.j/d&|
|d'| j  d%}d(|_0|st$j.1d)|t$j.1d*|t$j.1d+| j2t$j.1d,| j3rd-nd.t$j.1d/| j dkrd.nd-gn"t$j.1d)|t$j.1d*|t$j.1d/| j dkrd.nd-t$j.1d+| j2g}|j45| g }| j d0v rO| jrt
d1| j d t| j| j t$j%| jdd}| j  d2|j&_'t6|j&| j | j7s(t8||	}t
t9| d3d4d5 |D  d6 |j45t$j.1d7|j&t$j.1d8|	j&t$j.1d9t9|j&j:d:krI|j;ndg nt<|	j&}t
t9| d; |j4+t$j.1d8|	j& t$j.=d<t>j?d=d>g}t$j.=d?t>j?d-g}t$j.=d@t>j?d-g}t$j.=dAt>j?d-g}t$j.=dBt>j?d-g}t$j.=dCt>j@d-g}t$j.=dDt>j@d-g}|s|||||||gn||||g}| j*rt$j.=dt>j?|g}|+| | j,rt$j.=dt>j?d=|g}|+| | j-rt$j.=dt>j?d=d>g}|+| |st$j.=dt>j?g dEn
t$j.=dt>j?d=d?g}t$j.=d t>j@d=dBg}t$j.=d"t>j@dFd=dA|g}|g}| jr;|+| | jrD|+| t$j.A|g|sR| j  dGn| j  dH|||}t$j.jB|dI|	jCdJ} | jrd.dKlDmE}! |!Ft$jG|!FdLk rt
HdM tIjJ| | jdddN nt$J| | j t
dO| j  dPS )QzConvert model according to command line arguments.

    Args:
        args (argparse.Namespace): arguments parsed from command line
    r   z3Currently only gpt2 with greedy search is supportedzCoutput_sequences_scores currently is not supported in greedy searchz?output_token_scores currently is not supported in greedy searchz)skip convert_to_onnx since path existed: zgpt2_past_{}.onnxrm   rl   zConvert GPT model z	 to onnx z ...z,skip convert_to_onnx since paths specified: z and zConvert model z to onnx ...z Run symbolic shape inference on z. The file will be overwritten.r   r   zConfig=rP   Tr   z decoderr   
max_length
min_length	num_beamsnum_return_sequenceslength_penaltyrepetition_penaltyr   r  r  r  rJ   r#   rK   r   	sequencessequences_scoresz8--output_token_scores requires --output_sequences_scoresscores
BeamSearchBeamSearch_)inputsoutputsr   GreedySearchGreedySearch_zcom.microsofteos_token_idpad_token_idno_repeat_ngram_sizerI   rM   r   r   r   r9   zSymbolic shape inference on z encoder and decoder initz shared initializers (c                 S   s   g | ]}|j qS r*   )r   r   r*   r*   r+   r         z,convert_generation_model.<locals>.<listcomp>z*) in subgraphs are moved to the main graphr  r	  decoder_start_token_idr   z: initializers from the decoder are moved to the main graphr   
batch_sizesequence_lengthr  r  r  r  r  r  )r(  r  r  zmax_length - sequence_lengthz beam searchz greedy searchzonnxruntime.transformers)producer_nameopset_imports)versionz1.12.0z0Require onnx >= 1.12 to save large (>2GB) model!)r   all_tensors_to_one_filezmodel save to N)Kr   r$   r1   NotImplementedErrorrG   rH   rt   r\   r]   existsrx   ry   formatru   r	   r`   r   r   r   as_posixrs   r}   r   r   rE   r   rD   r   from_pretrainedr   r   r   r@   r"  r#  
vocab_sizer   r   r   r   r   r   rJ   rv   rK   rL   r   	make_nodedomainmake_attributer$  rI   	attributerw   r   rF   r  r   r   r'  r  r   r   r   r   
make_graph
make_modelopset_import	packagingr,  parse__version__r   r"   r   )"rj   r  is_gpt2is_greedysearchonnx_filenameconfigr"  r#  r3  r  r  r  r   attr_to_extendr
  r  r   r  r  r  r  r  r  graph_inputsrJ   rK   r   r  r  r  graph_outputs	new_graph	new_modelr,  r*   r*   r+   convert_generation_model  s  











	



	




	



	
rG  r   r   r   r"  r#  bad_words_idsc                 C   s   | j rtj std| jtjkr|  t	| j rdnd}|
| td |
|}|
|}g }t| jD ]/}	t }
|j||| j| j| j| j| j||| j| j| j|d| jp^| jd}	|t |
  q;|jd }ddlm} |||S )	a  Test PyTorch performance of text generation.

    Args:
        args (argparse.Namespace): arguments parsed from command line
        model (Union[GPT2LMHeadModel, T5ForConditionalGeneration]): PyTorch model
        input_ids (torch.Tensor): input_ids
        attention_mask (torch.Tensor): Attention mask
        eos_token_id (int): EOS token ID
        pad_token_id (int): Padding token ID
        bad_words_ids (List[List[int]]): Words shall not be generated.

    Raises:
        RuntimeError: PyTorch with CUDA is not available for --use_gpu

    Returns:
        Dict[str, Any]: A dictionary with string with metric name, and value can be integer or string.
    z=Please install PyTorch with Cuda for testing gpu performance.zcuda:0cpuFTr   r   r  r  r  rI   r$  r"  r#  r  r  r  rH  return_dict_in_generateoutput_scoresr   get_latency_result)rR   torchcudais_availabler   ru   r	   r`   halfdevicetoset_grad_enabledr   
total_runstimegenerater  r  r  rI   r$  r  r  r  rG   rH   rv   r   benchmark_helperrN  )rj   r   r   r   r"  r#  rH  rS  torch_latency_startr(  rN  r*   r*   r+   test_torch_performance4  sB   






r]  F	sentences	is_greedyc           +      C   s  | j dksJ tj| j| jd}d|_|j|_tj| j| j|j	d}|du r*g d}||ddd	}|d
 }|d }d}|j
|dd}	dd |	D }	| jrStd|	 ng }	|j}
|
j	}|
j	}|
j}g }d}| jstd td |j||| j| j| j| j| j||| j| j| j|	r|	ndd| jp| jd}td
| td td|j | jrtd|j | jrtd|j t |jD ]\}}|j!|dd}|"| t| d|  qtd td t#| j$| j%}|r|& ' (t)j*t)j+| jgt)j*dt)j+| jgt)j*dt)j+| jgt)j,dd}nB|& ' (t)j*t)j+| jgt)j*dt)j+| jgt)j*dt)j+| jgt)j*dt)j+| jgt)j*dt)j+| jgt)j,dt)j+| jgt)j,dd}| jrnt)j-|t)j*d}| jrj|	D ]}d||< qb||d< |j.d }| j/rt0d  t)j-||ft)j*d}||d!< td"| |1d|}| j2rt3| j$j45 }td#| dd$l6m7} |g}t |D ]\}}t8j9:|d%t;| }||| qg }t<| j=D ]}t>> }|1d|}|"t>> |  qdd&l?m@}  |j.d }| ||}!td' |d }"td|" | jrtd|d(  | jrtd|d)  |rG|"j.\}}#g }$t<|D ]}|j!|"| dd}|$"| td*| d+|  q*n5|"j.\}}%}#g }$t<|D ](}t<|%D ] }&|j!|"| |& dd}|$"| td*| d,|& d|  qYqS|r|jA|| jd-}'tBC|"}(td td. t|' t| td td/ t|( t|$ td ||$k})td0|)rd1nd2 |)|!d3< | jDrtE| ||||||	}*td4|* td5|! |!S )6a9  Test GPT-2 model

    Args:
        args (argparse.Namespace): arguments parsed from command line
        sentences (Optional[List[str]], optional): input text. Defaults to None.

    Returns:
        Union[Dict[str, Any], None]: A dictionary with string with metric name, and value can be integer or string.
    r   r  left)r   r#  N)zThe product is releasedzI enjoy walking in the parkzTest best way to investptTreturn_tensorspaddingr   r   walk in park)add_prefix_spacec                 S      g | ]}|gqS r*   r*   r   word_idr*   r*   r+   r     r&  z"test_gpt_model.<locals>.<listcomp>rH  2--------------------------------------------------CTest PyTorch model and beam search with huggingface transformers...rJ  !huggingface transformers outputs:r  r  r  skip_special_tokens: 'Testing beam search with onnxruntime...dtyper  r  r   rJ   zYUse prefix vocab mask with all ones in ORT, but no corresponding setting for Torch model.rK   
ORT inputstest_data_diroutput_test_datatest_data_set_rM  ORT outputs:rM   r   batch z sequence: 
 sequence rP   Torch Sequences:ORT Sequences:Torch and ORT result is same	differentparityTorch LatencyORT)Fr   r   r2  rs   r   padding_side	eos_token	pad_tokenr   r"  encoderJ   rx   r   rA  r3  rS   printrX  r  r  r  rI   r$  r  r  r  rG   rH   r  r  r  r   decoderv   r   r   rR   rI  numpyastypenpint32arrayfloat32onesr   rK   ry   runrU   r   r   r1  bert_test_datarv  r\   r]   r[   rZ   r   rV  rW  rY  rN  reshaperO  
LongTensorrT   r]  )+rj   r^  r_  	tokenizerr   r  r   r   	bad_wordsrH  rA  r"  r#  r3  torch_decoded_sequencesbeam_outputsr   sequencedecoded_sequencer   rJ   bad_word_idr(  rK   resultrt  rv  
all_inputsdirlatencyr[  r\  rN  r   r  r  ort_decoded_sequencesnum_sequencesr   torch_sequencesort_sequencesis_sametorch_latency_outputr*   r*   r+   test_gpt_modelv  s4  















	
r  c           *      C   s  | j dv sJ | jrtd dS tj| j| jd}d|_| j dkr,t	j| j| jd}n	t
j| j| jd}|du r=ddg}||d	d
d}|d }|d }d}||dd }dd |D }| jrhtd| ng }|j}	|	j}
|	j}|	j}td|
 d| d|  g }| jstd td |j||| j| j| j| j| j|
|| j| j| j|r|ndd
| jp| jd}td| td td|j | jrtd|j | jrtd|j  t!|jD ]\}}|j"|d
d}|#| td$|| qtd td t%| j&| j'}t(j)|t(j*d }| jr|D ]}d!||< q|+ , -t(j*t(j.| jgt(j*d t(j.| jgt(j*d t(j.| jgt(j*d t(j.| jgt(j*d t(j.| jgt(j/d t(j.| jgt(j/d d"}| jrc||d#< | j0rt(j)|j1t(j*d }t2|j1d! D ]*}d!}t2|j1d$ D ]}|| | |kr|d!krd!|| |< q|d$7 }qqw||d< | j3rt4| j&j56 }td%| d!d&l7m8} |g}t!|D ]\}}t9j:;|d't<| }||| qtd(| g }t2| j=D ]}t>> }|?d|}|#t>> |  q|j1d! }d!d)l@mA}  | ||}!td* |d! }"td|" | jr'td|d$  | jr2td|d+  |"j1\}}#}$g }%t2|D ](}t2|#D ] }|j"|"| | d
d}|%#| td,| d-| d.|  qDq>| js|jB|| jd}&tCD|"}'td td/ t|& t| td td0 t|' t|% td ||%k}(td1|(rd2nd3 |(|!d4< | jErtF| ||||
||})td5|) td6|! |!S )7a=  Test T5 or MT5 model

    Args:
        args (argparse.Namespace): arguments parsed from command line
        sentences (Optional[List[str]], optional): input text. Defaults to None.

    Returns:
        Union[Dict[str, Any], None]: A dictionary with string with metric name, and value can be integer or string.
    r%  zLSkipping parity test as prefix vocab mask is not implemented by Hugging FaceNr  r`  r   z4translate English to French: The product is releasedzsummarize: research continues to show that pets bring real health benefits to their owners.Having a dog around can lead to lower levels of stress for both adults and kids.ra  Trb  r   r   re  rP   c                 S   rg  r*   r*   rh  r*   r*   r+   r   l  r&  z!test_t5_model.<locals>.<listcomp>rH  zeos_token_id:z, pad_token_id:z, vocab_size:rj  rk  rJ  rl  r  r  r  rm  z{}: {}rp  rq  r   r  rJ   rM   rt  ru  rw  rs  rM  rx  r   ry  rz  ro  r{  r|  r}  r~  r  r  r  r  )Gr   rK   rx   r   r   r2  rs   r   r  r   r   r  rJ   rA  r"  r#  r3  rS   r  rX  r  r  r  rI   r$  r  r  r  rG   rH   r  r  r  r   r  rv   r0  r   r   rR   r  r  r  rI  r  r  r  r  rL   r   r   rU   r   r   r1  r  rv  r\   r]   r[   rZ   rV  rW  r  rY  rN  r  rO  r  rT   r]  )*rj   r^  r  r   r  r   r   r  rH  rA  r"  r#  r3  r  r  r   r  r  r   rJ   r  abs_posr   rt  rv  r  r  r  r[  r\  r  r(  rN  r   r  r  r  r  r  r  r  r  r*   r*   r+   test_t5_model>  s0  













	
r  c                 C   s<  t | }t|j |jdv rC|jr tj|js td|j |j	r2tj|j	s2td|j	 |jr8|j	r>|j	rB|jsBtdn|j
rJtd|jdkoS|jdk}|jdkrb|rbt|tj nt| td |jdv rwt||d	}nt|||d
}|r|jrtd|j d|j d |S td|j  |S )a/  Main entry function

    Args:
        argv (Optional[List[str]], optional): _description_. Defaults to None.
        sentences (Optional[List[str]], optional): input text. Defaults to None.

    Raises:
        ValueError: Path does not exist: --encoder_decoder_init_onnx
        ValueError: Path does not exist: --decoder_onnx
        ValueError: --decoder_onnx and --encoder_decoder_init_onnx are not used together for T5

    Returns:
        Union[Dict[str, Any], None]: A dictionary with string with metric name, and value can be integer or string.
    r%  z1Path does not exist: --encoder_decoder_init_onnx z$Path does not exist: --decoder_onnx zB--decoder_onnx shall use together with --encoder_decoder_init_onnxz>custom_attention_mask is only supported in t5 with beam searchrM   r   zstart testing model...)r^  )r^  r_  zOutput files: r5   z.datazOutput file: )rk   r   r@   r   r   r\   r]   r/  r   rt   rL   r.  r  r  rG  r$   r1   rx   ry   r  r  rD   r   )r2   r^  rj   r_  r  r*   r*   r+   r     s<   



r   __main__r'   )T)r   r   )r   )NF)NN)S__doc__rV   loggingr\   sysrW  enumr   pathlibr   typingr   r   r   r   r   r  r  r   rO  rY  r	   r
   r   r   transformersr   r   r   r   r   r   r   r   onnxruntimer   r   r   r   r]   rv   r[   dirname__file__gpt2_helperr   models.gpt2.convert_to_onnxr   rz   r   models.t5.convert_to_onnxr   r   models.t5.t5_helperr    r!   
onnx_modelr"   	getLoggerrx   r$   rZ   	Namespacerk   r}   r   boolr   r   r   r   r   ra   r   r  r  r0   rG  Tensorr]  r  r  r-   r*   r*   r*   r+   <module>   s   (  
 u)8MG
f
  )



"B I $G
6
