o
    ;ήcF                     @   s   d dl Z d dlZd dlmZmZ d dlZd dlZd dlmZm	Z	 d dlm
Z d dlmZ eeZdd Zdd
dZdddZdd Zg dZG dd dZ								dddZdddZdS )    N)DictList)helpernumpy_helper)onnx_pb)versionc                 C   s   dd | D S )z|
    Convert numpy float16 to python int.

    :param np_list: numpy float16 list
    :return int_list: python int list
    c                 S   s.   g | ]}t t|d dd ddqS )H   N   )intbinviewzfill).0_ r   G/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/transformers/float16.py
<listcomp>   s   . z%_npfloat16_to_int.<locals>.<listcomp>r   )np_listr   r   r   _npfloat16_to_int   s   r   "\o>     @c                 C   sz   dd }t |d| ||| } t || | d| | } t ||| td|| } t |td| | | | } t | S )a?  
    Convert float32 numpy array to float16 without changing sign or finiteness.
    Positive values less than min_positive_val are mapped to min_positive_val.
    Positive finite values greater than max_finite_val are mapped to max_finite_val.
    Similar for negative values. NaN, 0, inf, and -inf are unchanged.
    c                 S   s   t | |k ||k S N)nplogical_and)abcr   r   r   between(   s   z&convert_np_to_float16.<locals>.betweenr   infz-inf)r   wherefloatfloat16)np_arraymin_positive_valmax_finite_valr   r   r   r   convert_np_to_float16    s   
r&   c                 C   s   t | tjstdt|  | jtjjkrOtjj| _| jr9t	t
| j||}t|}|| jdd< g | jdd< | jrOt
j| jdd}t	|||}| | _| S )a  Convert tensor float to float16.

    Args:
        tensor (TensorProto): the tensor to convert.
        min_positive_val (float, optional): minimal positive value. Defaults to 1e-7.
        max_finite_val (float, optional): maximal finite value. Defaults to 1e4.

    Raises:
        ValueError: input type is not TensorProto.

    Returns:
        TensorProto: the converted tensor.
    5Expected input type is an ONNX TensorProto but got %sNfloat32dtype)
isinstance
onnx_protoTensorProto
ValueErrortype	data_typeFLOATFLOAT16
float_datar&   r   arrayr   
int32_dataraw_data
frombuffertobytes)tensorr$   r%   float16_dataint_listfloat32_listfloat16_listr   r   r   convert_tensor_float_to_float162   s   

r>   c                 C   s   t | j}t| j| j|S r   )r   to_arrayshaper   make_tensor_value_infonamer0   )r9   r@   r   r   r   make_value_info_from_tensorW   s   rC   )ArrayFeatureExtractor	BinarizerCastMapCategoryMapperDictVectorizerFeatureVectorizerImputerLabelEncoderLinearClassifierLinearRegressor
NormalizerOneHotEncoderSVMClassifierSVMRegressorScalerTreeEnsembleClassifierTreeEnsembleRegressorZipMapNonMaxSuppressionTopKRoiAlignResizeRangeCumSumMinMaxUpsamplec                   @   s0   e Zd ZdZdejfddZdejfddZdS )	InitializerTrackerz'Class for keeping track of initializer.initializerc                 C   s   || _ g | _g | _d S r   )r`   
fp32_nodes
fp16_nodes)selfr`   r   r   r   __init__~   s   
zInitializerTracker.__init__nodec                 C   s$   |r
| j | d S | j| d S r   )ra   appendrb   )rc   re   is_node_blockedr   r   r   add_node   s   zInitializerTracker.add_nodeN)	__name__
__module____qualname____doc__r,   r-   rd   	NodeProtorh   r   r   r   r   r_   {   s    r_   Fc           $         s  |dksJ d|t ttjjksJ dd}|s2ttjtdkr2z
ddl	m
}	 |	}W nw t| tjs@tdt|  |du rFt}|du rLg }t|}t|}td	| d
| d  d| d| d| d|  g }
g }g }|dur||| } |
|  i }t }t }dd | jjD }dd | jjD }t tr fdd|D } fdd|D }n sg }g }t| jjD ]U\}}|j|v rdt| }|||j< ||j dt| }| jj }|| ||_tjj |jj!_"t#j$d|jg|gd|dg}| jj%&| || || qt| jjD ]V\}}|j|v rpdt| }|||j< ||j dt| }| jj }|| ||_tjj |jj!_"t#j$d|g|jgd|dg}| jj%&| || || qi }|
rg }|
D ]Z}t|tjr||j t|tj'rL|j(D ]}|j)tjj*kr|j|vsJ t+|||j< q|j%D ]}|j|v rqt,t-|jD ]}|j| |v r||j|  |j|< qt,t-|jD ]}|j| |v r||j|  |j|< q|j.|v p|j|v }|jD ]}||v r|| /|| q|r || q|j.dkr>|j0D ]}|jdkr<|j1dkr<d|_1 nq)|j0D ]}|| qAqt|tj2r}||j3 |j4D ]}|| q\|j5t6|j5|| |j7D ]	}t6|||}qst|tj'rt89|j|j|jD ]F}|jj!j"tjj*kr|j|vrtjj |jj!_"|| |j:dr|jj;j"j!j"tjj*kr|j|vrtjj |jj;j"j!_"|| qq{|}
|
sw|< D ],\}} |s| j=rt6| j(||| _(|t>| j( | j?r|st@dA| j= q|D ]}!t,t-|!jD ]V}|!j| }|D ]K}"||"jkrk| jj }||" |!jd  t| }||_tjj*|jj!_"|!jd! t| }t#j$d|g|gd|dg}| jj%&| ||!j|<  nq!qt,t-|!jD ]V}|!j| }#|D ]K}"|#|"jkr| jj }||" |!jd" t| }||_tjj*|jj!_"|!jd# t| }t#j$d|g|#gd|dg}| jj%&| ||!j|<  nqqvq| S )$aY  Convert model tensor float type in the ONNX ModelProto input to tensor float16.

    Args:
        model (ModelProto): The ONNX model to convert.
        min_positive_val (float, optional): minimal positive value. Defaults to 5.96e-08.
        max_finite_val (float, optional): maximal finite value of float16. Defaults to 65504.
        keep_io_types (Union[bool, List[str]], optional): It could be boolean or a list of float32 input/output names.
                                                          If True, model inputs/outputs should be left as float32. Defaults to False.
        disable_shape_infer (bool, optional): Skips running onnx shape/type inference. Useful if shape inference has been done. Defaults to False.
        op_block_list (List[str], optional): List of op types to leave as float32.
                                             Defaults to None, which will use `float16.DEFAULT_OP_BLOCK_LIST` as default.
        node_block_list (List[str], optional): List of node names to leave as float32. Defaults to None.
        force_fp16_initializers(bool): force converting all float initializers to float16.
                                       Default to false, which will convert only the one needed to avoid precision loss.
    Raises:
        ValueError: input type is not ModelProto.

    Returns:
        ModelProto: converted model.
    r   zginvalid min_positive_val. smallest positive float16 value: subnormal 5.96e-08, and normalized 6.104e-05z4invalid max_finite_val. largest float16 value: 65504Nz1.2.0r   )infer_shapesz4Expected model type is an ONNX ModelProto but got %sz"fp16 parameters: min_positive_val=z max_finite_val=z keep_io_types=z disable_shape_infer=z op_block_list=z node_block_list=z force_fp16_initializers=c                 S   $   g | ]}|j jjtjjkr|jqS r   r/   tensor_type	elem_typer,   r-   r1   rB   r   nr   r   r   r         $ z,convert_float_to_float16.<locals>.<listcomp>c                 S   ro   r   rp   rs   r   r   r   r      ru   c                       g | ]}| v r|qS r   r   rs   keep_io_typesr   r   r          c                    rv   r   r   rs   rw   r   r   r      ry   graph_input_cast_graph_input_castCast
   )torB   graph_output_cast_graph_output_cast   r~   sequence_typezZinitializer is used by both fp32 and fp16 nodes. Consider add these nodes to block list:{}_input_cast__input_cast_output_cast__output_cast)Br!   r   finfor"   maxr   parseonnx__version__onnx.shape_inferencern   r+   r,   
ModelProtor.   r/   DEFAULT_OP_BLOCK_LISTsetloggerdebugrf   graphinputoutputlist	enumeraterB   stradd
value_infoCopyFromr-   r2   rq   rr   r   	make_nodere   extend
GraphProtor`   r0   r1   r_   rangelenop_typerh   	attributeiAttributeProtoggraphstr>   tensors	itertoolschainHasFieldr   itemsrb   rC   ra   infoformat)$modelr$   r%   rx   disable_shape_inferop_block_listnode_block_listforce_fp16_initializersfunc_infer_shapern   queuevalue_info_list	node_listname_mappinggraph_io_to_skipio_castsfp32_inputsfp32_outputsr   rt   output_name	node_namenew_value_infonew_node
input_namefp32_initializers
next_levelqrg   r   attrkeyvaluere   r   r   r   rw   r   convert_float_to_float16   s^  ,




















C





r   c                 C   s   t | tjstdt|  | jtjjkrtdd}| jr$t	| j}| j
r/tj| j
dd}|du r7tdt|||}tt|t| S )zSMeasure the maximum absolute difference after converting a float tensor to float16.r'   z#Expected tensor data type is float.Nr(   r)   zexternal data not loaded!)r+   r,   r-   r.   r/   r0   r1   r3   r   r4   r6   r7   RuntimeErrorr&   amaxabsr(   )r9   r$   r%   float32_datar:   r   r   r   float_to_float16_max_diffx  s   r   )r   r   )r   r   FFNNF)r   loggingtypingr   r   numpyr   r   r   r   r   r,   	packagingr   	getLoggerri   r   r   r&   r>   rC   r   r_   r   r   r   r   r   r   <module>   s2   



%
 o