o
    ;ήc>                     @   s  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mZmZmZ d dlZd dlZd dl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Z d dl m!Z! d d	l"m#Z# d d
l$m%Z% e &e'Z(g dZ)ej*dej+dej,diZ-G dd deZ.G dd deZ/G dd deZ0G dd deZ1G dd deZ2e1ddfe2ddfe0ddfdZ3G dd dZ4G d d! d!Z5dS )"    N)Path)DictListTupleUnion)
GPT2ConfigGPT2LMHeadModel	GPT2ModelTFGPT2Modelz..)	Precision)float_to_float16_max_diff)IOBindingHelper)	OnnxModel)torch_onnx_export)
distilgpt2gpt2zgpt2-mediumz
gpt2-largezgpt2-xlMb@?g?g      @c                       ,   e Zd ZdZ fddZ fddZ  ZS )GPT2ModelNoPastState2Here we wrap a class to disable past state output.c                       t  | d S Nsuper__init__selfconfig	__class__ W/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/models/gpt2/gpt2_helper.pyr   -      zGPT2ModelNoPastState.__init__c                    s   t  j|dddS )NF)	use_cachereturn_dict)r   forwardr   	input_idsr   r    r!   r%   0   s   zGPT2ModelNoPastState.forward__name__
__module____qualname____doc__r   r%   __classcell__r    r    r   r!   r   *       r   c                       r   )TFGPT2ModelNoPastStater   c                    s   d|_ t | d S )NF)r#   r   r   r   r   r    r!   r   7   s   zTFGPT2ModelNoPastState.__init__c                    s   t  j|ddS )NF)r#   )r   callr&   r   r    r!   r%   ;   r"   zTFGPT2ModelNoPastState.forwardr(   r    r    r   r!   r/   4   s    r/   c                       s8   e Zd ZdZ fddZedd Z fddZ  ZS )MyGPT2ModelzMHere we wrap a class for Onnx model conversion for GPT2Model with past state.c                    r   r   r   r   r   r    r!   r   B   r"   zMyGPT2Model.__init__c                 C   s   t | d d tst | d d trUt| d |kr$t| d d dks&J g }t|D ] }|tj| d | d d| d | d dfdd q,| d t|fS | S )N   r      )dim)	
isinstancetuplelistlenrangeappendtorchcat	unsqueeze)result	num_layerpresentir    r    r!   post_processE   s   $(*zMyGPT2Model.post_processc                    &   t  j||||dd}t|| jjS NF)position_idsattention_maskpast_key_valuesr$   r   r%   r1   rB   r   n_layerr   r'   rE   rF   pastr>   r   r    r!   r%   X   s   zMyGPT2Model.forward)	r)   r*   r+   r,   r   staticmethodrB   r%   r-   r    r    r   r!   r1   ?   s    
r1   c                       r   )MyGPT2LMHeadModelzSHere we wrap a class for Onnx model conversion for GPT2LMHeadModel with past state.c                    r   r   r   r   r   r    r!   r   f   r"   zMyGPT2LMHeadModel.__init__c                    rC   rD   rH   rJ   r   r    r!   r%   i   s   zMyGPT2LMHeadModel.forwardr(   r    r    r   r!   rM   c   r.   rM   c                       r   )MyGPT2LMHeadModel_NoPaddinga  Here we wrap a class for Onnx model conversion for GPT2LMHeadModel with past state and no padding.
    When you always use batch_size=1 in inference, there is no padding in inputs. In such case, position_ids
    and attention_mask need no be in inputs.
    c                    r   r   r   r   r   r    r!   r   {   r"   z$MyGPT2LMHeadModel_NoPadding.__init__c                    s"   t  j||dd}t|| jjS )NF)rG   r$   rH   )r   r'   rK   r>   r   r    r!   r%   ~   s   z#MyGPT2LMHeadModel_NoPadding.forwardr(   r    r    r   r!   rN   u   s    rN   logitsTF
last_state)r   GPT2LMHeadModel_NoPaddingr	   c                   @   s8   e Zd Zdd ZdefddZdefddZdd	 Zd
S )
Gpt2Inputsc                 C   s   || _ || _|| _|| _d S r   )r'   rE   rF   rK   )r   r'   rE   rF   rK   r    r    r!   r      s   
zGpt2Inputs.__init__returnc                 C   s0   dd | j | j| jfD }| jr|| j |S )Nc                 S   s   g | ]}|d ur|qS r   r    .0vr    r    r!   
<listcomp>       z&Gpt2Inputs.to_list.<locals>.<listcomp>)r'   rE   rF   rK   extend)r   
input_listr    r    r!   to_list   s   zGpt2Inputs.to_listc                 C   s"   t dd | j| j| j| jfD S )Nc                 s   s    | ]	}|d ur|V  qd S r   r    rT   r    r    r!   	<genexpr>   s    z&Gpt2Inputs.to_tuple.<locals>.<genexpr>)r6   r'   rE   rF   rK   )r   r    r    r!   to_tuple   s   "zGpt2Inputs.to_tuplec                 C   sT   d }| j d ur| j jtjkr| j jtjdn| j }dd | jD }t| j| j	||S )Ndtypec                 S   s   g | ]	}|j tjd qS )r^   )tor;   float32rU   pr    r    r!   rW      s    z&Gpt2Inputs.to_fp32.<locals>.<listcomp>)
rF   r_   r;   float16r`   ra   rK   rR   r'   rE   )r   rF   rK   r    r    r!   to_fp32   s   
zGpt2Inputs.to_fp32N)	r)   r*   r+   r   r   r[   r   r]   re   r    r    r    r!   rR      s
    rR   c                    @   s  e Zd ZdZedddejejejfdededededed	ed
edejde	de	de	dej
dej
dej
defddZe	dWdedededededeeee f fddZedd ZedXddZedXddZedYd!d"ZedZd$d%Zeddddejejejfd&ed'e	d(e	de	de	dej
dej
dej
fd)d*Ze		d[d+d,Zeg d-fd.ed/ee fd0d1Zed\d3ed4efd5d6Zed\d3ed4efd7d8Zed9d: Zed]d;d<Ze	2		d^d3ed=eeejf d>eeee f d4ed?e	d@e	fdAdBZ edCdD Z!edEdF Z"eddGdGdHdIddddejejejddfdJdKZ#eddLddddejejejdMdIdNfdOdPZ$ed_dQdRZ%edddg dSfdefdTdUZ&dVS )`
Gpt2HelperzEA helper class for Gpt2 model conversion, inference and verification.FT
batch_sizepast_sequence_lengthsequence_lengthnum_attention_headshidden_sizer?   
vocab_sizedevicerd   has_position_idshas_attention_maskinput_ids_dtypeposition_ids_dtypeattention_mask_dtyperS   c                    s   |rt jnt jd| ||t|| g fddt|D }t jd|d | |f| d}d}|
rT|| }t j| |g| d}|dkrTtd|d }d|dd|f< d}|	rv| 	d	d }|
|dk d |dd|df |}t||||S )
zCreate random inputs for GPT2 model.
        Returns torch tensors of input_ids, position_ids, attention_mask and a list of past state tensors.
        r3   c                    s$   g | ]}t j d d d qS )r_   rm   g       @      ?)r;   rand)rU   _rm   
float_type
past_shaper    r!   rW      s   $ z/Gpt2Helper.get_dummy_inputs.<locals>.<listcomp>r   r2   )lowhighsizer_   rm   Nrs   )r;   rd   ra   intr9   randintonesrandomlongcumsummasked_fill_r`   rR   )rg   rh   ri   rj   rk   r?   rl   rm   rd   rn   ro   rp   rq   rr   rK   r'   rF   total_sequence_lengthpadding_positionrE   r    rw   r!   get_dummy_inputs   s@   
zGpt2Helper.get_dummy_inputsr   r   model_classc                 C   s~   |j }|j}|j}|j}t| d }	| ||	dkr|n|g}
d| ||| t|| g}|	|
i}t|D ]
}||dt| < q2|S )zAReturns a dictionary with output name as key, and shape as value.r2   rO   r3   present_)rj   rk   num_hidden_layersrl   MODEL_CLASSESr~   r9   str)rg   rh   ri   r   r   rj   rk   r?   rl   output_namelast_state_shapepresent_state_shapeoutput_shapesrA   r    r    r!   get_output_shapes   s&   	
zGpt2Helper.get_output_shapesc                 C   sZ   |D ](}|| v s
J | | }t || | kr*tjt || |j|jd| |< qd S )Nrs   )numpyprodnelementr;   emptyr_   rm   )output_buffersr   keybufferr    r    r!   auto_increase_buffer_size  s   
z$Gpt2Helper.auto_increase_buffer_sizec                 C   sD   |rt jnt j}i }|  D ]\}}t jt|||d||< q|S )zpReturns a dictionary of output name as key, and 1D tensor as value. The tensor has enough space for given shape.rs   )r;   rd   ra   itemsr   r   r   )r   rm   
is_float16	data_typer   nameshaper    r    r!   get_output_buffers  s
   zGpt2Helper.get_output_buffersc                 C   sH   | d    }t||d  }|rt|t|d  S t|S )zGReturns the maximum difference between PyTorch and OnnxRuntime outputs.r   ư>)cpur   absamax)torch_outputsort_outputsrelativeexpected_outputsdiffr    r    r!   diff_outputs"  s
   
zGpt2Helper.diff_outputsMbP?c           
   	   K   s   t j|d | d    ||d}td|  |}t|d }t|D ])}t j|d|  | d |    ||d}td| d| d|  |oM|}q%|s`t| |}	t	d|	d	 |S )
zReturns True if torch and ORT outputs are close for given thresholds, and False otherwise.
        Note: need kwargs since Gpt2BeamSearchHelper.compare_outputs has an extra parameter model_class
        r   )rtolatolz9PyTorch and OnnxRuntime output 0 (last_state) are close: r2   zPyTorch and OnnxRuntime layer z state (present_z) are close:z@PyTorch and OnnxRuntime results are not all close: max_abs_diff=z.5f)
r   allcloser   loggerdebugr8   r9   rf   r   info)
r   r   r   r   kwargsis_closeis_all_close
num_layerslayermax_abs_diffr    r    r!   compare_outputs,  s"   "

zGpt2Helper.compare_outputsr   c                 C   s  d}d}g }g }t t|D ]}|| }|dkr| d n| d |d    }	tj||	|dd}
|tt|	|  |oA|
}t|		 rRt
d| d t|		 rbt
d| d t|	 rrt
d	| d t|	 rt
d	| d t||	 }t| |j}|d
|| dd| d|| ddt|	| d |dkrttj|dd|j}ttj|	dd|	j}t||}q|t|}|t||||fS )a  Compare outputs from PyTorch and OnnxRuntime

        Args:
            torch_outputs (Tuple[Torch.Tensor]): PyTorch model output
            ort_outputs (List[numpy.ndarray]): OnnxRuntime output
            atol (float, optional): Absolute tollerance. Defaults to 1e-06.

        Returns:
            is_all_close(bool): whether all elements are close.
            max_abs_diff(float): maximum absolute difference.
            messages(str): a list of debug message for each output
        TFr   r2   )r   r   zPyTorch output z has nanz has infzORT output zdiff=z.9fz index=z ort=z torch=N)axis)r9   r8   r   r   r   r:   r   r   isnananyr   r   isinffabsunravel_indexargmaxr   floatarray_equalindexmax)r   r   r   r   is_top1_matched	max_diffsmessagesrA   
ort_outputtorch_outputr   r   idxort_max_indextorch_max_indexmax_diff_output_indexr    r    r!   compare_outputs_v2G  sF   (0zGpt2Helper.compare_outputs_v2onnx_model_pathverboseuse_external_data_formatc
                 C   s  | j }
|
j}tjddd|
j|
j||
j|d|||||	d}| }t	  | | }W d   n1 s3w   Y  dd t
|D }dd t
|D }|d jd	 |
jks`|d jd	 |
jks`J |d jd	 |
jkrld
ndg| }dddd|d dddi}|D ]	}ddd||< q|D ]	}ddd||< qdg}|rddd|d< |d |rddd|d< |d || t|d	krt|d |ksJ td|jj d|jd j d|d j d|d d j  t|jjddd |rAt ;}tj|d}t|jjddd t| t||d|||ddd|d tj|dd} tj | |ddd W d   dS 1 s:w   Y  dS t| t||d|||ddd|d dS ) z1Export GPT-2 model with past state to ONNX model.r2   F)rg   rh   ri   rj   rk   r?   rl   rm   rd   rn   ro   rp   rq   rr   Nc                 S      g | ]}d | qS )past_r    rU   rA   r    r    r!   rW         z*Gpt2Helper.export_onnx.<locals>.<listcomp>c                 S   r   )r   r    r   r    r    r!   rW     r   r   r3   rO   rP   r'   rg   seq_len)r   r2   past_seq_len)r2      total_seq_lenrE   rF   zShapes: input_ids=z past=z output=z	 present=T)parentsexist_okz	gpt2.onnx   )
argsfexport_paramsinput_namesoutput_namesdynamic_axesopset_versiondo_constant_foldingr   r   )load_external_data)save_as_external_dataall_tensors_to_one_file)!r   rI   rf   r   rj   rk   rl   r[   r;   no_gradr9   r   r:   rY   r8   r   r   r'   rK   r   parentmkdirtempfileTemporaryDirectoryospathjoinr   r6   onnx
load_modelr   save)modelrm   r   r   r   rn   ro   rp   rq   rr   r   r?   dummy_inputsrZ   outputs
past_namespresent_namesr   r   r   r   tmp_dir_nametemp_onnx_model_pathr    r    r!   export_onnx}  s   

,"



 6
$
zGpt2Helper.export_onnxc              	   K   s~   ddl m} ddlm}	 |d}
|	| d||d|
dd}|r7|r%t| nd|vr-d|d< |jddd	i| ||| d
S )zHOptimize ONNX model with an option to convert it to use mixed precision.r   )FusionOptions)optimize_modelr   F)
model_type	num_headsrk   	opt_leveloptimization_optionsuse_gpukeep_io_typesuse_symbolic_shape_inferTNr    )fusion_optionsr   	optimizerr   rf   auto_mixed_precisionconvert_float_to_float16save_model_to_file)r   optimized_model_pathr   rj   rk   r   r  r   r   r   r   mr    r    r!   optimize_onnx  s&   
zGpt2Helper.optimize_onnx)AddLayerNormalizationFastGelu
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 GPT-2 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 ['Add', 'LayerNormalization', 'FastGelu']
        Returns:
            parameters(dict): a dictionary of parameters used in float16 conversion
        c                 S   s   g | ]}|j qS r    )op_type)rU   noder    r    r!   rW   )  s    z3Gpt2Helper.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 : r   z-Failed to find MatMul node for logits. Found z	 of node )r   r  node_block_listforce_fp16_initializersz!auto_mixed_precision parameters: r  Tr    )setnodes
differencer   r   graphoutputr   output_name_to_noder  inputget_initializerr   r   warningr  )r  r  op_full_setfp32_op_setfp16_op_setlogits_output_nameis_weight_fp16_precisionr  r  last_matmul_nodeinitializerr  max_diffr   r  
parametersr    r    r!   r    sH   




zGpt2Helper.auto_mixed_precisionr   inputs
total_runsc           	      C   s   t d |  }t  | | }W d   n1 sw   Y  |dkr)|S g }t   t|D ]}t }| | }|t |  q4W d   n1 sRw   Y  t	|d t
| }t dt|d ||fS )zfRun inference of PyTorch model, and returns average latency in ms when total_runs > 0 besides outputs.zstart pytorch_inferenceNr     zPyTorch inference time = {} ms.2f)r   r   re   r[   r;   r   r9   timer:   sumr8   format)	r   r'  r(  rZ   r   latencyrv   startaverage_latencyr    r    r!   pytorch_inference[  s$   



zGpt2Helper.pytorch_inferencec                 C   s"  t d dt|j  i}|jdur.t|jD ]\}}t|  |d| < q|jdur?t|j  |d< |j	durPt|j	  |d< | 
d|}|dkr\|S g }t|D ]}t }	| 
d|}|t |	  qbt|d t| }
t d	t|
d
 ||
fS )zcRun inference of ONNX model, and returns average latency in ms when total_runs > 0 besides outputs.zstart onnxruntime_inferencer'   Nr   rF   rE   r   r)  z"OnnxRuntime Inference time = {} msr*  )r   r   r   ascontiguousarrayr'   r   rK   	enumeraterF   rE   runr9   r+  r:   r,  r8   r-  )ort_sessionr'  r(  
ort_inputsrA   past_ir   r.  rv   r/  r0  r    r    r!   onnxruntime_inferenceu  s(   



z Gpt2Helper.onnxruntime_inferencec              	   C   s   t | ||||||S )z)Returnas IO binding object for a session.)r   prepare_io_binding)r5  r'   rE   rF   rK   r   r   r    r    r!   r9    s   zGpt2Helper.prepare_io_bindingc                 C   s   t | |||S )z3Copy results to cpu. Returns a list of numpy array.)r   "get_outputs_from_io_binding_buffer)r5  r   r   return_numpyr    r    r!   r:    s   z-Gpt2Helper.get_outputs_from_io_binding_bufferr   r   r;  include_copy_output_latencyc              	   C   s   t d t| |j|j|j|j||}| | t	| |||}|dkr'|S g }	t
|D ]}
t }| | |rBt	| |||}
|	t |  q-t|	d t|	 }t dt|d ||fS )zUInference with IO binding. Returns outputs, and optional latency when total_runs > 0.z*start onnxruntime_inference_with_binded_ior   r)  z2OnnxRuntime with IO binding inference time = {} msr*  )r   r   rf   r9  r'   rE   rF   rK   run_with_iobindingr:  r9   r+  r:   r,  r8   r-  )r5  r'  r   r   r(  r;  r<  
io_bindingr   r.  rv   r/  r0  r    r    r!   $onnxruntime_inference_with_binded_io  s8   


z/Gpt2Helper.onnxruntime_inference_with_binded_ioc                 C   s   t d|  dd}t|| W d    n1 sw   Y  td|  d t d|  dd}t|| W d    n1 sBw   Y  td|  d d S )Nort_outputs_.picklewbz$ORT output are saved to ort_outputs_torch_outputs_z(Torch output are saved to torch_outputs_openpickledumpr   r   )rA   r   r   r   r    r    r!   save_outputs  s   zGpt2Helper.save_outputsc                 C   sT   t d|  dd}t|| W d    n1 sw   Y  td|  d d S )Ndummy_inputs_rA  rB  z!inputs are saved to dummy_inputs_rD  )rA   r   r   r   r   r    r    r!   save_inputs  s   zGpt2Helper.save_inputsr   i'  r2   c           +         s  |j }td| d d| d| d|	 d| d d}d	}d
}d}|r5t|||||	}t|||}d}d}g  dg| }| }t|D ]}t| }t	d|}t	d|}t	d|}t
d| d| d tj||||j|j|j|j|||
||||d} t|| }!|rt| | }"nt|||||	}#t| | ||#}"tj|!|"|d\}$}%}&}'}(t|%s |% |$r|d7 }|(r|d7 }||  d7  < |r|$std| d| d| d| d|% 
 t|'D ]\}})td| d|  | j d|)  q|r#t|%s|%d| kr#t||  t||"|! qH r1 fdddD }*ndd dD }*|d | |*d < fd!d"|D |*d#< |d | |*d$< |t  d | |*d%< td&| d'| d(|t   d)|  |d*| krtd+t|d | d,d- |*S ).zKGenerate random inputs and compare the results of PyTorch and Onnx Runtime.zRunning parity test (atol=z, test_cases=z, runs=z, use_io_binding=z, model_class=z, is_float16=z) ...      r3   Nr   r2   z#Running parity test for batch_size=z past_sequence_length=z...rp   rq   rr   )r   z
test_case=z batch_size=z sequence_length=z	 MaxDiff=	z: Name=z, d   c              	      s&   i | ]}d | d t |qS )max_diff_percentile_z{:.5f})r-  r   
percentilerb   )max_abs_diff_listr    r!   
<dictcomp>e  s    z*Gpt2Helper.test_parity.<locals>.<dictcomp>)2   Z   _   c   c                 S   s   i | ]}d | dqS )rP  nanr    rb   r    r    r!   rS  j  rX   rt   top1_match_ratec                    s   g | ]}|d    qS )rt   r    )rU   x)test_cases_per_runr    r!   rW   m  rX   z*Gpt2Helper.test_parity.<locals>.<listcomp>top1_match_rate_per_rundiff_pass_ratenan_ratezParity Test Cases=z	; Passed=z; Nan=z; Top1_Matched=gffffff?zParity is good: passed rate=z.0f%)r   r   r   rf   r   r   r9   r~   r   r   r   r   rj   rk   rI   rl   r1  r8  r?  r   r   r   r:   r3  get_outputsr   rJ  rH  r8   )+r5  r   rm   r   r   r   r[  r(  use_io_bindingr   rn   ro   rp   rq   rr   r   enable_pickle_outputr   max_batch_sizemax_past_seq_lenmax_seq_lenr   max_output_shapespassed_test_casestop1_matched_casestop1_matched_cases_per_runtotal_test_casesrA   run_idri   rh   rg   r   r   r   r   r   r   r   r   r   messager>   r    )rR  r[  r!   test_parity  s   (




 ( 
" zGpt2Helper.test_parityrO  rK      c                 C   s   |j }d}|rt|||||}t|||}tj||||j|j|j|j|||||	|
|d}|r;t	| ||\}}|S t
| ||||\}}|S )zCGenerate random inputs and measure average latency of Onnx Runtime.NrM  )r   rf   r   r   r   rj   rk   rI   rl   r8  r?  )r5  r   rm   r   r(  ra  r   rn   ro   rp   rq   rr   rg   ri   rh   r   r   r   r   rv   r.  r    r    r!   test_performancez  s<   

zGpt2Helper.test_performancec                 C   s:   t jddd|j|j|j|j|d||d }tj	| |S )zJIT trace for TorchScript.r2   F)rg   rh   ri   rj   rk   r?   rl   rm   rd   rn   ro   )
rf   r   rj   rk   rI   rl   r[   r;   jittrace)r   r   rm   rn   ro   rZ   r    r    r!   torchscript  s    zGpt2Helper.torchscriptrawfp32fp16int8c                 C   s  |}t j|rt|jd }n|dd  |dkr!|d| 7 }|r'|d7 }|rdddd	d
}d
D ]P}t j| |||  }	t j|	r||v rwzt	|	 t
d|	  W q2 tyv }
 zt
d|	 d|
j  W Y d}
~
q2d}
~
ww t
d| d|	  q2t jt j| ||d t jt j| |d |d t jt j| |d |d t jt j| |d	 |d d
S t j| |d t j| |d t j| |d t j| |d d
S )z=Build a  path name for given model based on given attributes.r}   /r   rv   _past _fp32_fp16_int8rs  zRemoved the existed directory: zFailed to remove the directory r  NzDirectory for z
 existed: z.onnxz
_fp32.onnxz
_fp16.onnxz
_int8.onnx)r   r   isdirr   partssplitr   existsshutilrmtreer   r   OSErrorstrerror)
output_dirmodel_name_or_pathr   has_past
new_folderremove_existing
model_namesuffixr   new_direr    r    r!   get_onnx_paths  sT   

$zGpt2Helper.get_onnx_pathsN)r   )F)r   r   )r   )FF)r   )T)r   TF)TT)'r)   r*   r+   r,   rL   r;   int32r~   rm   boolr_   rR   r   r   r   r   r   r   r   r   r   r   r   r   r	  r   r  r1  r8  r9  r:  Tensorr?  rH  rJ  rm  ro  rr  r  r    r    r    r!   rf      sV   
	
:"
		5	
w&>
2
	
 6rf   )6loggingr   rF  r   r  sysr   r+  pathlibr   typingr   r   r   r   r   r   r;   transformersr   r   r	   r
   r   r:   r   dirname__file__benchmark_helperr   rd   r   io_binding_helperr   r  r   torch_onnx_export_helperr   	getLoggerr)   r   PRETRAINED_GPT2_MODELSFLOAT32FLOAT16INT8DEFAULT_TOLERANCEr   r/   r1   rM   rN   r   rR   rf   r    r    r    r!   <module>   sJ    

$