§
    kŠtjY  ã                   óò  — d dl Z d dlZd dlZd dlmZ d dlZd dlmZm	Z	m
Z
 d dlmZ de	dededed	ej        f
d
„Zde	deded	ej        fd„Zdedefd„Z	 d.de	dededededed	ej        fd„Zdedeeej        f         fd„Zdedededededede	de	de	dededefd„Z	 d/dededededede	de	de	dedededefd„Zd„ Z	 	 	 d0d ed!edz  d"edz  d#edz  d	eej        dz  ej        dz  ej        dz  f         f
d$„Z	 	 	 d0d%ed!edz  d"edz  d#edz  d	eej        dz  ej        dz  ej        dz  f         f
d&„Zd'„ Zd(ed)edededededed!edz  d"edz  d#edz  d*edededefd+„Zd,„ Ze d-k    r e¦   «          dS dS )1é    N)ÚPath)Ú
ModelProtoÚTensorProtoÚnumpy_helper)Ú	OnnxModelÚ	input_idsÚ
batch_sizeÚsequence_lengthÚdictionary_sizeÚreturnc                 óœ  — | j         j        j        t          j        t          j        t          j        fv sJ ‚t          j         	                    |||ft          j
        ¬¦  «        }| j         j        j        t          j        k    rt          j        |¦  «        }n3| j         j        j        t          j        k    rt          j        |¦  «        }|S )a`  Create input tensor based on the graph input of input_ids

    Args:
        input_ids (TensorProto): graph input of the input_ids input tensor
        batch_size (int): batch size
        sequence_length (int): sequence length
        dictionary_size (int): vocabulary size of dictionary

    Returns:
        np.ndarray: the input tensor created
    )ÚsizeÚdtype)ÚtypeÚtensor_typeÚ	elem_typer   ÚFLOATÚINT32ÚINT64ÚnpÚrandomÚrandintÚint32Úfloat32Úint64)r   r	   r
   r   Údatas        úe/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/onnxruntime/transformers/bert_test_data.pyÚfake_input_ids_datar      s°   € ð Œ>Ô%Ô/ÝÔÝÔÝÔð4ð ð ð ð õ Œ9×Ò˜_°JÀÐ3PÕXZÔX`ÐÑaÔa€Dà„~Ô!Ô+­{Ô/@Ò@Ð@ÝŒz˜$ÑÔˆˆØ	ŒÔ	#Ô	-µÔ1BÒ	BÐ	BÝŒx˜‰~Œ~ˆà€Kó    Úsegment_idsc                 ó„  — | j         j        j        t          j        t          j        t          j        fv sJ ‚t          j        ||ft          j	        ¬¦  «        }| j         j        j        t          j        k    rt          j
        |¦  «        }n3| j         j        j        t          j        k    rt          j        |¦  «        }|S )a,  Create input tensor based on the graph input of segment_ids

    Args:
        segment_ids (TensorProto): graph input of the token_type_ids input tensor
        batch_size (int): batch size
        sequence_length (int): sequence length

    Returns:
        np.ndarray: the input tensor created
    ©r   )r   r   r   r   r   r   r   r   Úzerosr   r   r   )r    r	   r
   r   s       r   Úfake_segment_ids_datar$   1   s©   € ð ÔÔ'Ô1ÝÔÝÔÝÔð6ð ð ð ð õ Œ8�Z Ð1½¼ÐBÑBÔB€DàÔÔ#Ô-µÔ1BÒBÐBÝŒz˜$ÑÔˆˆØ	Ô	Ô	%Ô	/µ;Ô3DÒ	DÐ	DÝŒx˜‰~Œ~ˆà€Kr   Úmax_sequence_lengthÚaverage_sequence_lengthc                 óœ   — |dk    r|| k    sJ ‚d|z  | k    rt          j        d|z  | z
  | ¦  «        S t          j        dd|z  dz
  ¦  «        S )Né   é   )r   r   )r%   r&   s     r   Úget_random_lengthr*   L   so   € Ø" aÒ'Ð'Ð,CÐGZÒ,ZÐ,ZÐ,ZÐZð 	Ð"Ñ"Ð%8Ò8Ð8ÝŒ~˜aÐ"9Ñ9Ð<OÑOÐQdÑeÔeÐeåŒ~˜a Ð%<Ñ!<¸qÑ!@ÑAÔAÐAr   r)   Ú
input_maskÚrandom_sequence_lengthÚ	mask_typec                 óz  — | j         j        j        t          j        t          j        t          j        fv sJ ‚|dk    rbt          j        |t          j	        ¬¦  «        }|r't          |¦  «        D ]}t          ||¦  «        ||<   Œ�nÎt          |¦  «        D ]}|||<   Œ�nµ|dk    r¦t          j        ||ft          j	        ¬¦  «        }|r=t          |¦  «        D ]+}t          ||¦  «        }t          |¦  «        D ]	}	d|||	f<   Œ
Œ,�nNt          j        ||ft          j	        ¬¦  «        }
|
|d|
j        d         …d|
j        d         …f<   �n	|dk    sJ ‚t          j        |dz  dz   t          j	        ¬¦  «        }|r‘t          |¦  «        D ]}t          ||¦  «        ||<   Œt          |dz   ¦  «        D ]X}|dk    r|||z   dz
           ||dz
           z   nd|||z   <   |dk    r|||z   dz
           ||dz
           z   nd|d|z  dz   |z   <   ŒYnHt          |¦  «        D ]}|||<   Œt          |dz   ¦  «        D ]}||z  |||z   <   ||z  |d|z  dz   |z   <   Œ| j         j        j        t          j        k    rt          j        |¦  «        }n3| j         j        j        t          j        k    rt          j        |¦  «        }|S )a"  Create input tensor based on the graph input of segment_ids.

    Args:
        input_mask (TensorProto): graph input of the attention mask input tensor
        batch_size (int): batch size
        sequence_length (int): sequence length
        average_sequence_length (int): average sequence length excluding paddings
        random_sequence_length (bool): whether use uniform random number for sequence length
        mask_type (int): mask type - 1: mask index (sequence length excluding paddings). Shape is (batch_size).
                                     2: 2D attention mask. Shape is (batch_size, sequence_length).
                                     3: key len, cumulated lengths of query and key. Shape is (3 * batch_size + 2).

    Returns:
        np.ndarray: the input tensor created
    r(   r"   r)   Nr   é   )r   r   r   r   r   r   r   r   Úonesr   Úranger*   r#   Úshaper   r   )r+   r	   r
   r&   r,   r-   r   ÚiÚactual_seq_lenÚjÚtemps              r   Úfake_input_mask_datar7   V   sV  € ð0 Œ?Ô&Ô0ÝÔÝÔÝÔð5ð ð ð ð ð �A‚~€~ÝŒw˜
­2¬8Ð4Ñ4Ô4ˆØ!ð 	2Ý˜:Ñ&Ô&ð Vð V�Ý+¨OÐ=TÑUÔU��Q‘�ñVõ ˜:Ñ&Ô&ð 2ð 2�Ø1��Q‘�ñ2à	�aŠˆÝŒx˜ _Ð5½R¼XÐFÑFÔFˆØ!ð 	:Ý˜:Ñ&Ô&ð #ð #�Ý!2°?ÐD[Ñ!\Ô!\�Ý˜~Ñ.Ô.ð #ð #�AØ!"�D˜˜A˜‘J�Jð#ñ#õ
 ”7˜JÐ(?Ð@ÍÌÐQÑQÔQˆDØ59ˆD��4”:˜a”=� / D¤J¨q¤M /Ð1Ñ2Ñ2à˜AŠ~ˆ~ˆ~ˆ~ÝŒx˜ a™¨!Ñ+µB´HÐ=Ñ=Ô=ˆØ!ð 	KÝ˜:Ñ&Ô&ð Vð V�Ý+¨OÐ=TÑUÔU��Q‘�å˜:¨™>Ñ*Ô*ð fð f�ØQRÐUVÒQVÐQV t¨J¸©N¸QÑ,>Ô'?À$ÀqÈ1ÁuÄ+Ñ'MÐ'MÐ\]��Z !‘^Ñ$ØYZÐ]^ÒY^ÐY^¨t°JÀ±NÀQÑ4FÔ/GÈ$ÈqÐSTÉuÌ+Ñ/UÐ/UÐde��Q˜‘^ aÑ'¨!Ñ+Ñ,Ð,ðfõ ˜:Ñ&Ô&ð 2ð 2�Ø1��Q‘�Ý˜:¨™>Ñ*Ô*ð Kð K�Ø'(Ð+BÑ'B��Z !‘^Ñ$Ø/0Ð3JÑ/J��Q˜‘^ aÑ'¨!Ñ+Ñ,Ð,à„Ô"Ô,µÔ0AÒAÐAÝŒz˜$ÑÔˆˆØ	ŒÔ	$Ô	.µ+Ô2CÒ	CÐ	CÝŒx˜‰~Œ~ˆà€Kr   Ú	directoryÚinputsc           	      ób  — t           j                             | ¦  «        sL	 t          j        | ¦  «         t	          d| › d�¦  «         n6# t
          $ r t	          d| › d�¦  «         Y nw xY wt	          d| › d�¦  «         t          |                     ¦   «         ¦  «        D ]Ž\  }\  }}t          j	        ||¦  «        }t          t           j                             | d|› d�¦  «        d	¦  «        5 }|                     |                     ¦   «         ¦  «         d
d
d
¦  «         n# 1 swxY w Y   Œ�d
S )z²Output input tensors of test data to a directory

    Args:
        directory (str): path of a directory
        inputs (Dict[str, np.ndarray]): map from input name to value
    z#Successfully created the directory ú zCreation of the directory z failedzWarning: directory z$ existed. Files will be overwritten.Úinput_ú.pbÚwbN)ÚosÚpathÚexistsÚmkdirÚprintÚOSErrorÚ	enumerateÚitemsr   Ú
from_arrayÚopenÚjoinÚwriteÚSerializeToString)r8   r9   ÚindexÚnamer   ÚtensorÚfiles          r   Úoutput_test_datarP   Ÿ   sŽ  € õ Œ7�>Š>˜)Ñ$Ô$ð Uð	FÝŒH�YÑÔÐõ ÐD¸	ÐDÐDÐDÑEÔEÐEÐEøõ ð 	Cð 	Cð 	CÝÐA¨yÐAÐAÐAÑBÔBÐBÐBÐBð	Cøøøõ
 	ÐS IÐSÐSÐSÑTÔTÐTå(¨¯ª©¬Ñ8Ô8ð 3ð 3Ñˆ‰|��dÝÔ(¨¨tÑ4Ô4ˆÝ•"”'—,’,˜yÐ*=°5Ð*=Ð*=Ð*=Ñ>Ô>ÀÑEÔEð 	3ÈØ�JŠJ�v×/Ò/Ñ1Ô1Ñ2Ô2Ð2ð	3ð 	3ð 	3ñ 	3ô 	3ð 	3ð 	3ð 	3ð 	3ð 	3ð 	3øøøð 	3ð 	3ð 	3ð 	3øð3ð 3s#   ¡A	 Á	A)Á(A)Ã/(D#Ä#D'	Ä*D'	Ú
test_casesÚverboseÚrandom_seedc           	      ó¸  — |€J ‚t           j                             |¦  «         t          j        |¦  «         g }t          |¦  «        D ]�}t	          || ||¦  «        }|j        |i}|rt          || |¦  «        ||j        <   |rt          || ||	|
|¦  «        ||j        <   |r#t          |¦  «        dk    rt          d|¦  «         | 
                    |¦  «         Œ‘|S )a  Create given number of input data for testing

    Args:
        batch_size (int): batch size
        sequence_length (int): sequence length
        test_cases (int): number of test cases
        dictionary_size (int): vocabulary size of dictionary for input_ids
        verbose (bool): print more information or not
        random_seed (int): random seed
        input_ids (TensorProto): graph input of input IDs
        segment_ids (TensorProto): graph input of token type IDs
        input_mask (TensorProto): graph input of attention mask
        average_sequence_length (int): average sequence length excluding paddings
        random_sequence_length (bool): whether use uniform random number for sequence length
        mask_type (int): mask type 1 is mask index; 2 is 2D mask; 3 is key len, cumulated lengths of query and key

    Returns:
        List[Dict[str,numpy.ndarray]]: list of test cases, where each test case is a dictionary
                                       with input name as key and a tensor as value
    Nr   zExample inputs)r   r   Úseedr1   r   rM   r$   r7   ÚlenrC   Úappend)r	   r
   rQ   r   rR   rS   r   r    r+   r&   r,   r-   Ú
all_inputsÚ
_test_caseÚinput_1r9   s                   r   Úfake_test_datar[   ¶   s	  € ðD Ð Ð Ð å„I‡N‚N�;ÑÔÐÝ
„K�ÑÔÐà€JÝ˜JÑ'Ô'ð "ð "ˆ
Ý% i°¸_ÈoÑ^Ô^ˆØ”. 'Ð*ˆàð 	gÝ'<¸[È*ÐVeÑ'fÔ'fˆF�;Ô#Ñ$àð 	Ý&:Ø˜J¨Ð9PÐRhÐjsñ'ô 'ˆF�:”?Ñ#ð ð 	,•s˜:‘”¨!Ò+Ð+ÝÐ" FÑ+Ô+Ð+Ø×Ò˜&Ñ!Ô!Ð!Ð!ØÐr   é'  rU   c                 ó~   — t          | ||||||||||	|
¦  «        }t          |¦  «        |k    rt          d¦  «         |S )aµ  Create given number of input data for testing

    Args:
        batch_size (int): batch size
        sequence_length (int): sequence length
        test_cases (int): number of test cases
        seed (int): random seed
        verbose (bool): print more information or not
        input_ids (TensorProto): graph input of input IDs
        segment_ids (TensorProto): graph input of token type IDs
        input_mask (TensorProto): graph input of attention mask
        average_sequence_length (int): average sequence length excluding paddings
        random_sequence_length (bool): whether use uniform random number for sequence length
        mask_type (int): mask type 1 is mask index; 2 is 2D mask; 3 is key len, cumulated lengths of query and key

    Returns:
        List[Dict[str,numpy.ndarray]]: list of test cases, where each test case is a dictionary
                                       with input name as key and a tensor as value
    z$Failed to create test data for test.)r[   rV   rC   )r	   r
   rQ   rU   rR   r   r    r+   r&   r,   r-   r   rX   s                r   Úgenerate_test_datar^   ð   s`   € õB  ØØØØØØØØØØØØñô €Jõ ˆ:�„˜*Ò$Ð$ÝÐ4Ñ5Ô5Ð5ØÐr   c                 ó  — |t          |j        ¦  «        k    rd S |j        |         }|                      |¦  «        }|€C|                      ||¦  «        }|�+|j        dk    r |                      |j        d         ¦  «        }|S )NÚCastr   )rV   ÚinputÚfind_graph_inputÚ
get_parentÚop_type)Ú
onnx_modelÚ
embed_nodeÚinput_indexra   Úgraph_inputÚparent_nodes         r   Úget_graph_input_from_embed_noderj   $  sŒ   € Ø•c˜*Ô*Ñ+Ô+Ò+Ð+ØˆtàÔ˜[Ô)€EØ×-Ò-¨eÑ4Ô4€KØÐØ ×+Ò+¨J¸ÑDÔDˆØÐ" {Ô':¸fÒ'DÐ'DØ$×5Ò5°kÔ6GÈÔ6JÑKÔKˆKØÐr   re   Úinput_ids_nameÚsegment_ids_nameÚinput_mask_namec                 ó  — |                       ¦   «         }|�Í|                      |¦  «        }|€t          d|› �¦  «        ‚d}|r)|                      |¦  «        }|€t          d|› �¦  «        ‚d}|r)|                      |¦  «        }|€t          d|› �¦  «        ‚d|rdndz   |rdndz   }t          |¦  «        |k    r"t          d|› dt          |¦  «        › �¦  «        ‚|||fS t          |¦  «        dk    rt          dt          |¦  «        › �¦  «        ‚|                      d	¦  «        }	t          |	¦  «        dk    rw|	d         }
t          | |
d¦  «        }t          | |
d¦  «        }t          | |
d
¦  «        }|€$|D ]!}|j                             ¦   «         }d|v r|}Œ"|€t          d¦  «        ‚|||fS d}d}d}|D ]/}|j                             ¦   «         }d|v r|}Œ"d|v sd|v r|}Œ-|}Œ0|r	|r|r|||fS t          d¦  «        ‚)a  Find graph inputs for BERT model.
    First, we will deduce inputs from EmbedLayerNormalization node.
    If not found, we will guess the meaning of graph inputs based on naming.

    Args:
        onnx_model (OnnxModel): onnx model object
        input_ids_name (str, optional): Name of graph input for input IDs. Defaults to None.
        segment_ids_name (str, optional): Name of graph input for segment IDs. Defaults to None.
        input_mask_name (str, optional): Name of graph input for attention mask. Defaults to None.

    Raises:
        ValueError: Graph does not have input named of input_ids_name or segment_ids_name or input_mask_name
        ValueError: Expected graph input number does not match with specified input_ids_name, segment_ids_name
                    and input_mask_name

    Returns:
        Tuple[Optional[np.ndarray], Optional[np.ndarray], Optional[np.ndarray]]: input tensors of input_ids,
                                                                                 segment_ids and input_mask
    Nz Graph does not have input named r(   r   zExpect the graph to have z inputs. Got r/   z'Expect the graph to have 3 inputs. Got ÚEmbedLayerNormalizationé   Úmaskz#Failed to find attention mask inputÚtokenÚsegmentz?Fail to assign 3 inputs. You might try rename the graph inputs.)Ú'get_graph_inputs_excluding_initializersrb   Ú
ValueErrorrV   Úget_nodes_by_op_typerj   rM   Úlower)re   rk   rl   rm   Úgraph_inputsr   r    r+   Úexpected_inputsÚembed_nodesrf   ra   Úinput_name_lowers                r   Úfind_bert_inputsr|   1  sÄ  € ð4 ×EÒEÑGÔG€LàÐ!Ø×/Ò/°Ñ?Ô?ˆ	ØÐÝÐPÀÐPÐPÑQÔQÐQàˆØð 	XØ$×5Ò5Ð6FÑGÔGˆKØÐ"Ý Ð!VÐDTÐ!VÐ!VÑWÔWÐWàˆ
Øð 	WØ#×4Ò4°_ÑEÔEˆJØÐ!Ý Ð!UÀOÐ!UÐ!UÑVÔVÐVà KÐ6˜q˜q°QÑ7À
Ð;Q¸1¸1ÐPQÑRˆÝˆ|ÑÔ Ò/Ð/ÝÐj¸ÐjÐjÕWZÐ[gÑWhÔWhÐjÐjÑkÔkÐkà˜+ zÐ1Ð1å
ˆ<ÑÔ˜AÒÐÝÐVÅ3À|ÑCTÔCTÐVÐVÑWÔWÐWà×1Ò1Ð2KÑLÔL€KÝ
ˆ;ÑÔ˜1ÒÐØ  ”^ˆ
Ý3°JÀ
ÈAÑNÔNˆ	Ý5°jÀ*ÈaÑPÔPˆÝ4°ZÀÈQÑOÔOˆ
àÐØ%ð 'ð '�Ø#(¤:×#3Ò#3Ñ#5Ô#5Ð ØÐ-Ð-Ð-Ø!&�JøØÐÝÐBÑCÔCÐCà˜+ zÐ1Ð1ð €IØ€KØ€JØð 	ð 	ˆØ œ:×+Ò+Ñ-Ô-ÐØÐ%Ð%Ð%ØˆJˆJàÐ'Ð'Ð'¨9Ð8HÐ+HÐ+HàˆKˆKàˆIˆIàð 2�[ð 2 Zð 2Ø˜+ zÐ1Ð1å
ÐVÑ
WÔ
WÐWr   Ú	onnx_filec                 óþ   — t          ¦   «         }t          | d¦  «        5 }|                     |                     ¦   «         ¦  «         ddd¦  «         n# 1 swxY w Y   t	          |¦  «        }t          ||||¦  «        S )aó  Find graph inputs for BERT model.
    First, we will deduce inputs from EmbedLayerNormalization node.
    If not found, we will guess the meaning of graph inputs based on naming.

    Args:
        onnx_file (str): onnx model path
        input_ids_name (str, optional): Name of graph input for input IDs. Defaults to None.
        segment_ids_name (str, optional): Name of graph input for segment IDs. Defaults to None.
        input_mask_name (str, optional): Name of graph input for attention mask. Defaults to None.

    Returns:
        Tuple[Optional[np.ndarray], Optional[np.ndarray], Optional[np.ndarray]]: input tensors of input_ids,
                                                                                 segment_ids and input_mask
    ÚrbN)r   rH   ÚParseFromStringÚreadr   r|   )r}   rk   rl   rm   ÚmodelrO   re   s          r   Úget_bert_inputsrƒ   �  s¯   € õ( ‰LŒL€EÝ	ˆi˜Ñ	Ô	ð + $Ø×Ò˜dŸiši™kœkÑ*Ô*Ð*ð+ð +ð +ñ +ô +ð +ð +ð +ð +ð +ð +øøøð +ð +ð +ð +õ ˜5Ñ!Ô!€JÝ˜J¨Ð8HÈ/ÑZÔZÐZs   Ÿ(AÁAÁAc                  ó  — t          j        ¦   «         } |                      ddt          d¬¦  «         |                      ddt          d d¬¦  «         |                      d	dt          d
d¬¦  «         |                      ddt          dd¬¦  «         |                      ddt          d d¬¦  «         |                      ddt          d d¬¦  «         |                      ddt          d d¬¦  «         |                      ddt          d
d¬¦  «         |                      ddt          dd¬¦  «         |                      dddd¬¦  «         |                      d¬¦  «         |                      dddd ¬¦  «         |                      d¬!¦  «         |                      d"d#d$t          d%¬&¦  «         |                      d'd(ddd)¬¦  «         |                      d¬*¦  «         |                      d+dt          d,d-¬¦  «         |                      ¦   «         }|S ).Nz--modelTzbert onnx model path.)Úrequiredr   Úhelpz--output_dirFz4output test data path. Default is current directory.)r…   r   Údefaultr†   z--batch_sizer(   zbatch size of inputz--sequence_lengthé€   z maximum sequence length of inputz--input_ids_namezinput name for input idsz--segment_ids_namezinput name for segment idsz--input_mask_namezinput name for attention maskz	--samplesz$number of test cases to be generatedz--seedr/   zrandom seedz	--verboseÚ
store_truezprint verbose information)r…   Úactionr†   )rR   z--only_input_tensorsz-only save input tensors and no output tensors)Úonly_input_tensorsz-az--average_sequence_lengthéÿÿÿÿz)average sequence length excluding padding)r‡   r   r†   z-rz--random_sequence_lengthz3use uniform random instead of fixed sequence length)r,   z--mask_typer)   z^mask type: (1: mask index, 2: raw 2D mask, 3: key lengths, cumulated lengths of query and key))ÚargparseÚArgumentParserÚadd_argumentÚstrÚintÚset_defaultsÚ
parse_args)ÚparserÚargss     r   Úparse_argumentsr–   ©  s¼  € ÝÔ$Ñ&Ô&€Fà
×Ò˜	¨DµsÐAXÐÑYÔYÐYà
×ÒØØÝØØCð ñ ô ð ð ×Ò˜°½SÈ!ÐRgÐÑhÔhÐhà
×ÒØØÝØØ/ð ñ ô ð ð ×ÒØØÝØØ'ð ñ ô ð ð ×ÒØØÝØØ)ð ñ ô ð ð ×ÒØØÝØØ,ð ñ ô ð ð ×ÒØØÝØØ3ð ñ ô ð ð ×Ò˜¨5µsÀAÈMÐÑZÔZÐZà
×ÒØØØØ(ð	 ñ ô ð ð ×Ò ÐÑ&Ô&Ð&à
×ÒØØØØ<ð	 ñ ô ð ð ×Ò¨5ÐÑ1Ô1Ð1à
×ÒØØ#ØÝØ8ð ñ ô ð ð ×ÒØØ"ØØØBð ñ ô ð ð ×Ò¨uÐÑ5Ô5Ð5à
×ÒØØÝØØmð ñ ô ð ð ×ÒÑÔ€DØ€Kr   r‚   Ú
output_dirr‹   c                 óÞ  — t          | |||	¦  «        \  }}}t          |||||||||||¦  «        }t          |¦  «        D ]E\  }}t          j                             |dt          |¦  «        z   ¦  «        }t          ||¦  «         ŒF|
rdS ddl}d| 	                    ¦   «         v rddgndg}| 
                    | |¬¦  «        }d„ |                     ¦   «         D ¦   «         }t          |¦  «        D ]þ\  }}t          j                             |dt          |¦  «        z   ¦  «        }|                     ||¦  «        }t          |¦  «        D ]£\  }}t          j        t          j        ||         ¦  «        |¦  «        }t#          t          j                             |d|› d	�¦  «        d
¦  «        5 }|                     |                     ¦   «         ¦  «         ddd¦  «         n# 1 swxY w Y   Œ¤ŒÿdS )aI  Create test data for a model, and save test data to a directory.

    Args:
        model (str): path of ONNX bert model
        output_dir (str): output directory
        batch_size (int): batch size
        sequence_length (int): sequence length
        test_cases (int): number of test cases
        seed (int): random seed
        verbose (bool): whether print more information
        input_ids_name (str): graph input name of input_ids
        segment_ids_name (str): graph input name of segment_ids
        input_mask_name (str): graph input name of input_mask
        only_input_tensors (bool): only save input tensors,
        average_sequence_length (int): average sequence length excluding paddings
        random_sequence_length (bool): whether use uniform random number for sequence length
        mask_type(int): mask type
    Útest_data_set_Nr   ÚCUDAExecutionProviderÚCPUExecutionProvider)Ú	providersc                 ó   — g | ]	}|j         ‘Œ
S © )rM   )Ú.0Úoutputs     r   ú
<listcomp>z-create_and_save_test_data.<locals>.<listcomp>N  s   € ÐDÐDÐD F�F”KÐDÐDÐDr   Úoutput_r=   r>   )rƒ   r^   rE   r?   r@   rI   r�   rP   ÚonnxruntimeÚget_available_providersÚInferenceSessionÚget_outputsÚrunr   rG   r   ÚasarrayrH   rJ   rK   )r‚   r—   r	   r
   rQ   rU   rR   rk   rl   rm   r‹   r&   r,   r-   r   r    r+   rX   r3   r9   r8   r£   rœ   ÚsessionÚoutput_namesÚresultÚoutput_nameÚtensor_resultrO   s                                r   Úcreate_and_save_test_datar®     s\  € õD *9¸ÀÐP`ÐbqÑ)rÔ)rÑ&€Iˆ{˜Jå#ØØØØØØØØØØØñô €Jõ ˜zÑ*Ô*ð ,ð ,‰	ˆˆ6Ý”G—L’L Ð-=ÅÀAÁÄÑ-FÑGÔGˆ	Ý˜ FÑ+Ô+Ð+Ð+àð ØˆàÐÐÐð # k×&IÒ&IÑ&KÔ&KÐKÐKð 
!Ð"8Ð9Ð9à$Ð%ð ð
 ×*Ò*¨5¸IÐ*ÑFÔF€GØDÐD¨g×.AÒ.AÑ.CÔ.CÐDÑDÔD€Lå˜zÑ*Ô*ð >ð >‰	ˆˆ6Ý”G—L’L Ð-=ÅÀAÁÄÑ-FÑGÔGˆ	Ø—’˜\¨6Ñ2Ô2ˆÝ'¨Ñ5Ô5ð 	>ð 	>‰NˆAˆ{Ý(Ô3µB´J¸vÀa¼yÑ4IÔ4IÈ;ÑWÔWˆMÝ•b”g—l’l 9Ð.>¸Ð.>Ð.>Ð.>Ñ?Ô?ÀÑFÔFð >È$Ø—
’
˜=×:Ò:Ñ<Ô<Ñ=Ô=Ð=ð>ð >ð >ñ >ô >ð >ð >ð >ð >ð >ð >øøøð >ð >ð >ð >øð	>ð>ð >s   Æ,(G Ç G$Ç'G$c                  ó>  — t          ¦   «         } | j        dk    r| j        | _        | j        }|€It	          | j        ¦  «        }t          j                             |j	        d| j
        › d| j        › �¦  «        }|�'t	          |¦  «        }|                     dd¬¦  «         nt          d¦  «         t          | j        || j
        | j        | j        | j        | j        | j        | j        | j        | j        | j        | j        | j        ¦  «         t          d|¦  «         d S )Nr   Úbatch_Ú_seq_T)ÚparentsÚexist_okz7Directory existed. test data files will be overwritten.z Test data is saved to directory:)r–   r&   r
   r—   r   r‚   r?   r@   rI   Úparentr	   rB   rC   r®   ÚsamplesrU   rR   rk   rl   rm   r‹   r,   r-   )r•   r—   Úpr@   s       r   Úmainr·   Y  s   € ÝÑÔ€DàÔ# qÒ(Ð(Ø'+Ô';ˆÔ$à”€JØÐå�”ÑÔˆÝ”W—\’\ !¤(Ð,a°T´_Ð,aÐ,aÈ4ÔK_Ð,aÐ,aÑbÔbˆ
àÐå�JÑÔˆØ�
Š
˜4¨$ˆ
Ñ/Ô/Ð/Ð/åÐGÑHÔHÐHåØŒ
ØØŒØÔØŒØŒ	ØŒØÔØÔØÔØÔØÔ$ØÔ#ØŒñô ð õ" 
Ð
,¨jÑ9Ô9Ð9Ð9Ð9r   Ú__main__)r)   )r\   )NNN)!r�   r?   r   Úpathlibr   Únumpyr   Úonnxr   r   r   re   r   r‘   Úndarrayr   r$   r*   Úboolr7   r�   ÚdictrP   r[   r^   rj   Útupler|   rƒ   r–   r®   r·   Ú__name__rž   r   r   ú<module>rÁ      sñ  ðð €€€Ø 	€	€	€	Ø €€€Ø Ð Ð Ð Ð Ð à Ð Ð Ð Ø 6Ð 6Ð 6Ð 6Ð 6Ð 6Ð 6Ð 6Ð 6Ð 6Ø  Ð  Ð  Ð  Ð  Ð  ðØðØ(+ðØ>AðØTWðà„Zðð ð ð ð< {ð Àð ÐVYð Ð^`Ô^hð ð ð ð ð6B¨3ð BÈð Bð Bð Bð Bð  ðFð FØðFàðFð ðFð !ð	Fð
 !ðFð ðFð „ZðFð Fð Fð FðR3 ð 3¨T°#°r´z°/Ô-Bð 3ð 3ð 3ð 3ð.7Øð7àð7ð ð7ð ð	7ð
 ð7ð ð7ð ð7ð ð7ð ð7ð !ð7ð !ð7ð ð7ð 7ð 7ð 7ðL !ð1ð 1Øð1àð1ð ð1ð ð	1ð
 ð1ð ð1ð ð1ð ð1ð !ð1ð !ð1ð ð1ð ð1ð 1ð 1ð 1ðh
ð 
ð 
ð "&Ø#'Ø"&ð	YXð YXØðYXà˜$‘JðYXð ˜D‘jðYXð ˜4‘Zð	YXð
 ˆ2Œ:˜Ñ˜bœj¨4Ñ/°´¸dÑ1BÐBÔCðYXð YXð YXð YXð| "&Ø#'Ø"&ð	[ð [Øð[à˜$‘Jð[ð ˜D‘jð[ð ˜4‘Zð	[ð
 ˆ2Œ:˜Ñ˜bœj¨4Ñ/°´¸dÑ1BÐBÔCð[ð [ð [ð [ð8að að aðHI>ØðI>àðI>ð ðI>ð ð	I>ð
 ðI>ð ðI>ð ðI>ð ˜$‘JðI>ð ˜D‘jðI>ð ˜4‘ZðI>ð ðI>ð !ðI>ð !ðI>ð ðI>ð I>ð I>ð I>ðX$:ð $:ð $:ðN ˆzÒÐØ€D�F„F€F€F€Fð Ðr   