o
    ;ήcJ                     @   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Zd dlZd dlm	Z	m
Z
 d dlmZmZmZ d dlmZ d dlmZ ejejejedd d dlmZmZmZmZ d d	lmZ ed
Z dddZ!dd Z"e#dkre! Z$ee$j% e"e$ dS dS )    N)datetime)MODEL_CLASSESGpt2HelperFactory)DEFAULT_TOLERANCEPRETRAINED_GPT2_MODELS
Gpt2Helper)version)
AutoConfigz..)	Precisioncreate_onnxruntime_sessionprepare_environmentsetup_logger)QuantizeHelper c                 C   s  t  }|jdddtddt d |jddtd	tt d
dt  d |jddtt	j
dddd |jddtt	j
dddd |jdddtdd |jdddddd |jdddddd |jdd  |jd!ddd"d |jdd# |jd$d%ttjttd&d' |jd(ddd)d |jdd* |jd+d,d-td.gd/d0 |jd1td2d3d4 |jd5d-td.gd6d0 |jd7d8d-tg d9d:d0 |jd;d<dd d=d> |jd?dtd@dAd |jdBdddC |jddD |jdEdddC |jddF |dG}|jdHtddId4 |jdJtd.dKd4 |jdLtd.dMd4 |jdNdd-tdOdP |jdQtd.dRd4 |dS}|jdTddUdV |jdWtdXdYd4 |jdZtd[d\d4 || }|S )]Nz-mz--model_name_or_pathTz;Model path, or pretrained model name selected in the list: z, )requiredtypehelpz--model_classFGPT2LMHeadModelz!Model type selected in the list: )r   r   defaultchoicesr   z--cache_dir.cache_modelsz%Directory to cache pre-trained models)r   r   r   r   z
--onnx_dironnx_modelszDirectory to store onnx modelsz--test_timesd   z8Number of repeat times to get average inference latency.)r   r   r   r   z-vz--validate_onnx
store_truezValidate ONNX model)r   actionr   z-oz--optimize_onnxz'Use optimizer.py to optimize onnx model)optimize_onnxz	--use_gpuzuse GPU for inference)use_gpuz-pz--precisionzfPrecision of model to run. fp32 for full precision, fp16 for half precision, and int8 for quantization)r   r   r   r   z--torchscriptzuse Torchscript)torchscriptz-bz--batch_sizes+   z
batch size)nargsr   r   r   z--beam_size   z2Beam size if greedy/top-p/top-k sampling is needed)r   r   r   z--sequence_lengthsz!sequence lengths (excluding past)z-sz--past_sequence_lengths)          @         zpast sequence lengthsz-rz--result_csvz$CSV file for saving summary results.)r   r   r   z--thread_numzThreads to usez--include_copy_output_latency)r   r   )include_copy_output_latencyz	--verbose)verbosez$configurable one step search optionsz--ignore_eosz3If ignore end of sentence token in model inference.z--repetition_penaltyz,Positive. >1 to penalize and <1 to encorage.z--temperaturez&Softmax temperature for output logits.z--excluded_token_idsz0A list of token ids to be excluded in inference.)r   r!   r   r   z--length_penaltyz;Positive. >1 to penalize and <1 to encorage short sentence.zone step sampling optionsz--do_samplez3If to do sampling instead of beam search or greedy.)r   r   z--do_sample_top_pgffffff?z0Nuclear/top-p sampling accumulation probability.z--do_sample_top_kr   zUse top-k if non-zero.)argparseArgumentParseradd_argumentstrjoinr   listr   keysospathintset_defaultsr
   FLOAT32add_argument_groupboolfloat
parse_args)argvparsersearch_option_groupsampling_option_groupargs rA   Z/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/gpt2/benchmark_gpt2.pyparse_arguments   s"  
			


rC   c           %      C   s   ddl m} t|tdk rtdtd|   | jtj	kr,| j
r(| js,J d| jtjkr9| jr9J dt| jdkrFtjdd	n| j ttj  | j}| j}t||| j t| j d }| jd
krmd}n
| jdkrud}nd}t|}tj| j| j|d}|dkr|j| j|d| j |d}n*|dkr|j| j|d| j | j!| j"| j#| j$| j%| j&| j'| j(|d}n	|j| j||d}t)| jrdnd}	|*|	 |j+dk}
|j,|| j| jd|
d}|d }t| j d }|j-||	|| j.|
||d | j
s| jtj/krM|| jtjkrt0| jnd }|j
|d || jtj	k|j1j2|j1j3|
dd | jtjkrMtd t45||d |
 t46|}td |d }| jr[|j|||	||d}t7|| jd | j| j.d!}|d u rnd S |dksx|dkr|j8t9| j:t9| j;t9| j;t9| j<| j d|| jd"}|=||	| jtj	k}n|8t9| j:t9| j;t9| j<|| j}|=||	| jtj	k}| j>pd#?t@A Bd$}tC|d%d&d'}g d(}tDjE||d)}|F  | j:D ]~}| j<D ]v}| j;D ]n}|dkr|dkr|dksJ tGd*| d+| d,| d- |dks|dkrF|jH||||j2|j3|j+|jI|	| jtj	k||d.}|8||||| j d|| j}n"|jH||||j2|j3|j+|jI|	| jtj	k||d.}|8||||| j}z|J||| jK\}}tL|D ],\}}tM|tNrtGd/| d0tO| d1|d jP  qwtGd/| d2|jP  qw|Q||| jK\}}|jR||||| jKd | jSd3\} }!| jTr|jU||| jtV| j tV| j d4rtd5tV| j  d6 g }"| D ]}#|"W|#X Y  q|jU||"| jtV| j tV| j d4rtd7tV| j  d6 td8| d9| d:| d;|d<d=|d<d>|!d< | j| j| j| j| j
| j||||d<|d<|!d<d(}$|Z|$ W q   tj[d?dd@ Y    W d    d S qqW d    n	1 sqw   Y  tdA|  |S )BNr   )__version__z3.1.0z/This tool requires transformers 3.1.0 or later.z
Arguments:z'fp16 requires --optimize_onnx --use_gpuzquantization only supports CPUT)logicalGPT2LMHeadModel_BeamSearchStepbeam_search_step)GPT2LMHeadModel_ConfigurableOneStepSearchconfigurable_one_step_searchr   )r   	cache_dirr    )config
batch_size	beam_sizerJ   )rK   rL   rM   
ignore_eostemperaturerepetition_penaltyexcluded_token_idslength_penalty	do_sampledo_sample_top_pdo_sample_top_krJ   )rK   rJ   zcuda:0cpu   )has_past
new_folderraw   )has_position_idshas_attention_maskfp32)auto_mixed_precisionzquantizing model...int8zfinished quantizing modelF)enable_all_optimizationnum_threadsr+   )context_lenpast_sequence_lengthsequence_lengthrM   steprK   model_classzbenchmark_result_{}.csvz%Y%m%d-%H%M%Sar   )modenewline)
model_namerg   gpu	precision	optimizerr   rL   re   rd   torch_latencyonnxruntime_latencyonnxruntime_io_binding_latency)
fieldnameszRunning test for batch_size=z sequence_length=z past_sequence_length=z...)float16r\   r]   ztorch output z is tuple of size z, shape z shape )return_numpyr*   )rg   rtolatolz:Pytorch and ONNX Runtime outputs are all close (tolerance=z).zEPytorch and ONNX Runtime IO Binding outputs are all close (tolerance=zbatch_size=z, sequence_length=z, past_sequence_length=z, torch_latency=z.2fz, onnxruntime_latency=z!, onnxruntime_io_binding_latency=	Exception)exc_infozResults are saved to file )\transformersrD   r   parseRuntimeErrorloggerinform   r
   FLOAT16r   r   INT8torchset_num_threads
thread_numpsutil	cpu_countprint
__config__parallel_inforJ   onnx_dirr   r   rg   r   create_helperr	   from_pretrainedmodel_name_or_pathr   rM   rN   rO   rP   rQ   rR   rS   rT   rU   deviceton_layerget_onnx_pathsexport_onnxr+   r7   r/   rK   num_attention_headshidden_sizer   quantize_onnx_modelquantize_torch_modelr   get_output_shapesmaxbatch_sizespast_sequence_lengthssequence_lengthsget_output_buffers
result_csvformatr   nowstrftimeopencsv
DictWriterwriteheaderdebugget_dummy_inputs
vocab_sizepytorch_inference
test_times	enumerate
isinstancetuplelenshapeonnxruntime_inference$onnxruntime_inference_with_binded_ior*   validate_onnxcompare_outputsr   appendrV   numpywriterowerror)%r@   transformers_versionrJ   
output_dirrg   
model_type
gpt2helperrK   modelr   use_external_data_formatonnx_model_pathsonnx_model_pathuse_paddingsessionmax_output_shapesoutput_bufferscsv_filenamecsv_filecolumn_names
csv_writerrL   re   rd   dummy_inputsoutput_shapesoutputsro   ivalueort_outputsort_latencyort_io_outputsort_io_latencycopy_outputsoutputrowrA   rA   rB   main   s  "





 





"

*

,  r   __main__)N)&r,   r   loggingr3   sysr   r   r   gpt2_beamsearch_helperr   r   gpt2_helperr   r   r   	packagingr   ry   r	   r4   r   r0   dirname__file__benchmark_helperr
   r   r   r   quantize_helperr   	getLoggerr|   rC   r   __name__r@   r+   rA   rA   rA   rB   <module>   s4    

 +  .
