o
    ;ήcRs                     @   sL  d dl Z d dlZd dlZd dl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 d dlZd dlZd dlmZmZmZ d dlmZ d dlZejejejedd d dlZe dZ!d	eee"ef  fd
dZ#d5ddZ$							d6d	eee"ef  fddZ%d	ee"ef fddZ&dd Z'd7ddZ(d	ee"ef fddZ)d	eee"ef  fddZ*d	eee"ef  fddZ+d	eee"ef  fd d!Z,d8d"d#Z-d$d% Z.d	eee"ef  fd&d'Z/d	eee"ef  fd(d)Z0								*d9d+d,Z1d:d.d/Z2d;d0d1Z3d2d3 Z4e5d4kr$e4  dS dS )<    N)ProcessPoolExecutor)datetime)AnyDictList)PRETRAINED_LONGFORMER_MODELSLongformerHelperLongformerInputs)LongformerModelz.. returnc                    s:  |dkr	t | g }|D ]}	|D ]}
|D ]}td|	 d|
 d|  t|	|
|| }|    }tj fdd|dd}i d	d
dt j	ddddddddd|d|d ddd|d|	d|
d|dt
t ddddd ddddd!}|t||	 td"| || qqq|S )#Nr   batch_size= sequence_length= global_length=c                      s     S N r   
input_listmodelr   f/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/longformer/benchmark_longformer.py<lambda>S   s    z$test_torch_latency.<locals>.<lambda>   )repeatnumberenginetorchversiondevicecuda	optimizerr   	precisionfp32
io_binding
model_namedescriptionz [torch]inputs   threads
batch_sizesequence_lengthglobal_lengthr   memoryNAdiff_maxdiff_90_percentile)diff_95_percentilediff_99_percentileuse_compact_memory%s)r   set_num_threadsloggerinfor   get_dummy_inputsto_listtimeitr   __version__strr   nowupdatebenchmark_helperget_latency_resultappend)r   r   r#   batch_sizessequence_lengthsglobal_lengths
test_timesnum_threadsresultsr(   r)   r*   r%   _runtimesresultr   r   r   test_torch_latency;   st   

	
#rI   Tc                 C   s   d| d| d| }t d| d t|||| }| }	|d |	}
| }|| }t|d 	 
 |
d  }t d|  |r^t|sP|dkr^td	|d  td
|
d  t|S )Nr   r   r   z$Comparing Torch and ORT outputs for ...r   zlast_state max diff = gMbP?ztorch last_state:zort last_state:)r4   r5   r   r6   get_ort_inputsrunr7   npamaxcpunumpymathisnanprintfloat)r   r   ort_sessionr(   r)   r*   verbose
parametersdummy_inputs
ort_inputsort_outputsr   torch_outputsmax_diffr   r   r   test_parityp   s   r]   Fr!   c                    s   g }|D ] |D ]|D ]j jd ksJ dtd  d d d|
 d| d|  d	 t }| }rGt| d |}i d
|d|dddddt	t
jdddt	|dt|
dt|	dt dtdtdt|dt	t dddd dd d d ||d}|st|t| j j }t|t| }tj|||d d!g|g ||g tjd"}n
tj||| d#}|s fd$d%t|D }t||d< t|d&|d< t|d'|d(< t|d)|d*< || qq	q|S )+Nr   zPLimitation of current implementation: number of global token <= attention_windowzTesting batch_size=r   r   z optimizer=z, precision=z io_binding=rJ   r#   r$   r%   r&   r   OnnxRuntimer   r   r   r    r   r'   r(   r)   r*   rC   r   r+   r   r-   r.   )r/   r0   r1   	use_half4
last_statepooler)	result_templaterepeat_timesort_output_namesrZ   output_buffersoutput_buffer_max_sizesr(   r   	data_type)rb   rc   r(   c              
      s    g | ]}t  qS r   )r]   ).0rF   r(   r   r*   r   rU   r)   rV   r   r   
<listcomp>   s    
z$test_ort_latency.<locals>.<listcomp>Z   _   r/   c   r0   )configattention_windowr4   r5   r   r6   rK   rS   rL   r:   onnxruntimer9   intr   r;   maxhidden_sizer=   inference_ort_with_io_bindingrM   longlonginference_ortrange
percentiler?   )r   r   r#   r$   rU   r@   rA   rB   rC   rD   r   r    disable_io_bindingrV   r1   r_   disable_parityrE   rX   rY   rZ   rb   max_last_state_sizemax_pooler_sizerH   diff_resultsr   ri   r   test_ort_latency   s   


	
	
]r~   c           	         sh   t d d  d d d d   fdd}tjd	|d
} |dS )NzTesting memory for model=z, batch_size=z, sequence_length=z, global_length=z, test_times=z, num_threads=c                     sZ   ddi} d| i}t jdd|d}t }| }tD ]}|d |}q"d S )Narena_extend_strategykSameAsRequestedCUDAExecutionProviderTuse_gpuenable_all_optimizationrD   provider_options)r=   create_onnxruntime_sessionr   r6   rK   rw   rL   )cuda_provider_optionsr   sessionrX   rY   rF   r(   r   r*   rD   onnx_model_pathr)   rC   r   r   	inference  s    z"test_ort_memory.<locals>.inferenceT)is_gpufunc)
onnx_modelr(   r)   r*   rC   rD   r+   )r4   r5   r=   measure_memory)	r   r   r(   r)   r*   rC   rD   r   memory_usedr   r   r   test_ort_memory   s&   	r   c                 C   s,   | t v rt |  n| }t|}|| |S r   )r   r
   from_pretrainedto)r#   r   torch_model_name_or_dirr   r   r   r   load_torch_model%  s
   

r   .c                 C   s^   t j|| d }t j|| d }t j|| d }t j|r%|}|S t j|r-|}|S )Nz.onnx
_fp32.onnx
_fp16.onnx)ospathjoinisfile)r#   onnx_dirr   optimized_fp32_modeloptimized_fp16_modelr   r   r   find_onnx_model.  s   r   c                 C   s   t | jdkrtdt | jdkrtdt | jdkr!td| j}| js+t|n| j}tj	
  t||| jd | jd | jd | j| jS )Nr   z5For memory test, only one batch_size (-b) is allowed.z:For memory test, only one sequence_length (-s) is allowed.z8For memory test, only one global_length (-g) is allowed.r   )lenr@   RuntimeErrorrA   rB   r   onnxr   r   r   empty_cacher   rC   rD   )argsr   r#   r   r   r   r   test_memory:  s$   
r   c                 C   s  | j }| js
t|n| j}|dp|d}|dsdnd}t||}| j}ddi}d|i}	tj|dd||	d	}
|
d u rEtd
| t	j
dddk}|}|sV|d7 }| jrd||dkr`dnd7 }n
||dkrkdnd7 }t|||||
| j| j| j| j|||| j| j|| j| jS )Nr   r   r!   fp16r   r   r   Tr   z,Failed to create ORT session from ONNX file ORT_LONGFORMER_COMPACT_MEMORY1z[non_compact_memory]z[half4]z[float4]z[half2])r   r   r   endswithr   rD   r=   r   r   r   environgetr_   r~   r@   rA   rB   rC   ry   rV   rz   )r   r   r#   r   	optimizedr    r   rD   r   r   r   r1   r$   r   r   r   test_ortQ  sV   
r   c              	   C   s.   t | j|}t||| j| j| j| j| j| jS r   )r   r   rI   r@   rA   rB   rC   rD   )r   r   r   r   r   r   
test_torch  s   r   c                 C   s   d| j kr
t| |S t| |S )Nrp   )r   r   r   r   r   r   r   r   test_latency  s   


r   c                 C   s\  t  }|jdddtdddt  d |jdd	dtd
d
dgdd |jddddtdd |jdddtdgd |jdddtg ddd |jddtd dd |jdd dtd!gd"d |jd#d$dtd!d%d |jd&dd'd(d) |jd*dd'd+d) |jd,dd'd-d) |jdd. |jd/dd'd0d) |jdd1 |jd2dd'd3d) |jdd4 |	| }|S )5Nz-mz--modelFlongformer-base-4096z=Checkpoint directory or pre-trained model names in the list: z, )requiredtypedefaulthelpz-ez--enginerp   r   zEngine to benchmark.)r   r   r   choicesr   z-tz--test_times  z8Number of repeat times to get average inference latency.)r   r   r   r   z-bz--batch_sizes+r   )nargsr   r   z-sz--sequence_lengths)            zSequence lengths. It could have multiple values in latency test.If --export_padding is not used, sequence length shall be multiple of window size.)r   r   r   r   z--onnxzOnnx model pathz-gz--global_lengthsr   zGNumber of global tokens. It could have multiple values in latency test.z-nz--num_threadszThreads to use.z--disable_io_binding
store_truezDo not use IO Binding.)r   actionr   z--memoryz%Test memory usage instead of latency.z	--verbosezPrint more information.)rV   z--use_half4zUse half4 kernel.)r_   z--disable_parityzDo not run parity test.)rz   )
argparseArgumentParseradd_argumentr:   r   r   keysrq   set_defaults
parse_args)argvparserr   r   r   r   parse_arguments  s   

	
		
r   c                 C   s   dd | D }t |dkrtd d S t|dddd)}g d	}tj||d
}|  |D ]}t| || q-|  W d    n1 sGw   Y  td|  d S )Nc                 S   s   g | ]}d |v r|qS average_latency_msr   rh   rH   r   r   r   rj     s    z"output_details.<locals>.<listcomp>r   zNo latency results for output.ar   asciimodenewlineencoding)r   r   r   r    r   r"   r#   r%   r'   r   rC   r$   r(   r)   r*   r1   r_   r-   r.   r/   r0   r+   QPSr   latency_variancelatency_90_percentilelatency_95_percentilelatency_99_percentile
fieldnamesz&Detail results are saved to csv file: )r   rS   opencsv
DictWriterwriteheaderwriterowflush)rE   csv_filenamelatency_resultscsv_filecolumn_names
csv_writerrH   r   r   r   output_details  s   
(r   c                 C   s:   t d td t d}| jrt| |gS t| |S )NF{   zcuda:0)r   set_grad_enabledr=   set_random_seedr   r+   r   r   r   r   r   r   rL   -  s   



rL   c                 C   sf   t j s	tdt }t|t| g}t|dksJ |d W  d    S 1 s,w   Y  d S )NzYPlease install PyTorch with Cuda, and use a machine with GPU for testing gpu performance.r   r   )	r   r   is_availabler   r   listmaprL   r   )	argumentsexecutorrE   r   r   r   launch_test<  s   
$r   r   c                 C   sB  | rdnd}|t jd< td|  |rdndt jd< td|r$dnd g }	d}
g d}|g}d	D ]}|D ]}|D ]}d
D ]}|rid}td| d|
 d| d| d| d|
 d| d}|	t|7 }	d}|rodnd}|r{| d| dn| d| d}t j	|st
d| d| d| d| d| d| d| }|s|d7 }|r|d7 }|dkr|d7 }d }z"|rt| dd}t|}t| d|
 d}t|}W n ty } zt
d |d }~w   t  Y qAt|dkr|r|d d! nd"|d d!< nt
d#td$| |	|7 }	qAq=q9q5|	S )%Nr   0r   zORT_LONGFORMER_COMPACT_MEMORY=ORT_LONGFORMER_USE_HALF4zORT_LONGFORMER_USE_HALF4={}r   )r   r   r   r   )r   )   r   z-e z -t z -b z -s z -g z -m  rp   r   r   _fr   r   zonnx file not exists:z --onnx z --disable_io_bindingz --use_half4   z --disable_parityz -t 10 --memoryzKeyboard Interruptedr+   zN/Az%length of latency_results should be 1r2   )r   r   r4   r5   formatr   splitrL   r   existsr   r   KeyboardInterrupt	traceback	print_excr   )r1   	run_torch
run_memoryuse_io_bindinguse_fp16use_merged_qkv_weightsr_   r(   compact_memoryrE   rC   rA   r@   r#   r)   r*   engine_namer   file_format	onnx_pathr   memory_resultsr   excr   r   r   	run_testsF  s   

 
 
:r  r   c                    s@  t |dddd}g d ttdd | D }|  ttdd | D }|  ttd	d | D }|  g }|D ]}|D ]}	|d
|	 d|  qCq?tj| | d}
|
  |D ]}i }i }|dd |D  i }|dd |D  | D ]b}|d |kr|| r fdd|	 D }|s|| n D ]}|| || krt
dq|d }	|d }d
|	 d| }zt|| }W n	 ty   Y qw ||  |7  < ||  d7  < q|r|D ]}||v r|| dkr|| ||  ||< qd||< q|
| qa|  W d    d S 1 sw   Y  d S )Nr   r   r   r   )r#   r    r   r   r*   r1   r_   r$   c                 S      g | ]}|d  qS )r$   r   r   r   r   r   rj         z"output_summary.<locals>.<listcomp>c                 S   r  )r(   r   r   r   r   r   rj     r  c                 S   r  )r)   r   r   r   r   r   rj     r  b_sr   c                 S      i | ]}|d qS r   r   rh   kr   r   r   
<dictcomp>      z"output_summary.<locals>.<dictcomp>c                 S   r  r  r   r  r   r   r   r    r  r$   c                    s   i | ]\}}| v r||qS r   r   )rh   r  vheader_namesr   r   r    s    zDescription shall be uniquer(   r)   r   r   )r   r   setsortr?   r   r   r   r<   itemsr   rT   
ValueErrorr   r   )rE   r   
data_fieldr   description_listr@   rA   
data_namesr)   r(   r   r$   rowsum_latencycount_latencyrH   headersr  keylatencyr   r  r   output_summary  sd   


$r!  c                 C   s\   t | dd|d}|r|S | r"|t | dd|d7 }|t | dd|d7 }|t | dd|d7 }|S )zARun experiments to compare different algorithms on one batch sizeTF)r   r   r_   r(   )r  )r   r(   is_baselinetest_resultsr   r   r   run_experiments  s8   r$  c                  C   s  t jd t } t| j ttj	dkr.t
| }t d}d| d}t|| d S t }td| g d}g d}|rR|d	 d
 dkrRg d}g d}|r_tdd|d	 d nd}tjdddk}d| |rrdnd }	td|	 d| d|  d}
g }t|
D ]}|D ]}td||d}t|d ||7 }qqdD ]}t||	 d| d| qg }t|
D ]}|D ]}td||d}t|d ||7 }qqdD ]}t||	 d| d| qd S ) Nspawnr   z%Y%m%d-%H%M%Sbenchmark_detail_z.csvzGPU info: %s)r      r      r   )r   r(  r   r   totall         )@       r   r'  r   r(  r   z(?u)[^-\w.]rF   namegpuORT_LONGFORMER_BASELINEr   r   longformer_base_	_baseliner   zexperiment_name=z, fp16_batch_sizes=z, fp32_batch_sizes=T)r   r(   r"  zlongformer_base_fp16.csv)r   r   r+   r.   Fzlongformer_base_fp32.csv)r   multiprocessingset_start_methodr   r=   setup_loggerrV   r   sysr   r   r   r;   strftimer   get_gpu_infor4   r5   resubr   r   r   rw   r$  r!  )r   r#  
time_stampr   gpu_listfp16_batch_sizesfp32_batch_sizesgpu_namer"  experiment_name
total_runsall_resultsrF   r(   fp16_resultsmetric_namefp32_resultsr   r   r   main  sT   




rD  __main__)T)Fr!   FTFFF)r   r   )TFTTTTTr   r   )F)6r   r   loggingrQ   r   r7  r4  r8   r   concurrent.futuresr   r   typingr   r   r   rP   rM   r   longformer_helperr   r   r	   transformersr
   rp   r   r?   r   dirname__file__r=   	getLoggerr4   r:   rI   r]   r~   r   r   r   r   r   r   r   r   r   rL   r   r  r!  r$  rD  __name__r   r   r   r   <module>   s|     
	

5
u

.
	7
_1

S
I%
1
