o
    ;ήcg                 	   @   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mZmZ d dl	m
Z
 e
eje
dks1J eeZd+ddZdd Zd	d
 Zdd Zdd Zdd Zdd Zdd Zdd Zd+ddZdd Zdd Zdd ZG dd  d Zd!d" Zed#kre Ze d$ej!  ej"re d%ej"  e d& e#e$ej!ej%ej&ej'ej(Z)ej"re)rej*rej+e)ej"d'ej,ej-ej.d(d) ne/e)ej" e d* dS dS dS dS ),    N)helpernumpy_helpershape_inference)versionz1.8.0c                    s*    fdd| j D }|rt|d S |S )Nc                    s   g | ]	}|j  kr|qS  name).0attr	attr_namer   M/tmp/pip-target-vg8gfxp4/lib/python/onnxruntime/tools/symbolic_shape_infer.py
<listcomp>       z!get_attribute.<locals>.<listcomp>r   )	attributer   get_attribute_value)noder   default_valuefoundr   r   r   get_attribute   s   r   c                 C   s&   t | dtkrt| | dS d S )Nvalue)type
WhichOneofstrgetattrdimr   r   r   get_dim_from_proto   s   &r   c                 C   s   |  d}|dv sJ |dkS )Nr   )tensor_typesequence_typer   )r   )
type_protocls_typer   r   r   is_sequence   s   
r"   c                 C   s0   t | rJ | jdrdd | jjjD S d S )Nshapec                 S      g | ]}t |qS r   )r   r	   dr   r   r   r   '       z-get_shape_from_type_proto.<locals>.<listcomp>)r"   r   HasFieldr#   r   )r    r   r   r   get_shape_from_type_proto$   s   r)   c                 C   sR   | j d}|d u rd S t| j r$d| j jjdkr"t| j jjS d S t| j S )Nr   r   )r   r   r"   r   	elem_typer)   )vir!   r   r   r   get_shape_from_value_info,   s   

r,   c                 C   s   t  }| |_|S N)onnxValueInfoProtor   )r   r+   r   r   r   make_named_value_info9   s   r0   c                 C   s   dd | D S )Nc                 S   s0   g | ]}|d u r
d nt |rt|nt|qS r-   )
is_literalintr   r	   ir   r   r   r   @      0 z.get_shape_from_sympy_shape.<locals>.<listcomp>r   )sympy_shaper   r   r   get_shape_from_sympy_shape?      r7   c                 C   s*   t | ttjtjtjfv pt| do| jS )N	is_number)	r   r2   npint64int32sympyIntegerhasattrr9   r   r   r   r   r1   C   s   *r1   c                 C   s*   | |k r	| | ksJ | dkr| S ||  S Nr   r   )axisrankr   r   r   handle_negative_axisG   s   rC   c                 C   sB   |pg d}t |tkr|g}| jD ]}|j|v r|j  S qd S )N) r.   zai.onnx)r   listopset_importdomainr   )mprG   opsetr   r   r   	get_opsetL   s   


rJ   c                 C   s>   t | tkrt| dksJ | d S t | tjkr|  S | S N   r   )r   rE   lenr:   ndarrayitemxr   r   r   	as_scalarW   s   rR   c                 C   s<   t | tkr| S t | tjkrt| S |r| d u rd S | gS r-   )r   rE   r:   rN   )rQ   	keep_noner   r   r   as_lista   s   rT   c                 C   s4   t | tkrtd}| D ]}|| }q|S | }|S )NrL   )r   rE   r=   r>   )rQ   r   vr   r   r   sympy_reduce_productl   s   

rV   c                   @   s  e Zd ZdddZdddZdddZd	d
 Zdd Zdd Zdd Z	dd Z
dd Zdd Zdd Zdd Zdd ZdddZdd d!Zd"d# Zd$d% Zd&d' Zd(d) Zdd+d,Zdd-d.Zd/d0 Zdd1d2Zdd4d5Zd6d7 Zd8d9 Zd:d; Zd<d= Zd>d? Zd@dA Z dBdC Z!dDdE Z"dFdG Z#dHdI Z$dJdK Z%dLdM Z&dNdO Z'dPdQ Z(dRdS Z)dTdU Z*dVdW Z+dXdY Z,dZd[ Z-d\d] Z.d^d_ Z/d`da Z0dbdc Z1ddde Z2dfdg Z3dhdi Z4djdk Z5dldm Z6dndo Z7dpdq Z8drds Z9dtdu Z:dvdw Z;dxdy Z<dzd{ Z=d|d} Z>d~d Z?dd Z@dd ZAdd ZBdd ZCdd ZDdd ZEdd ZFdd ZGdd ZHdd ZIdd ZJdd ZKdd ZLdd ZMdd ZNdd ZOdd ZPdd ZQdd ZRdd ZSdd ZTdd ZUdd ZVdd ZWdd ZXdd ZYdd ZZdddZ[dd Z\dd Z]dddZ^dd Z_e`dddZad3S )SymbolicShapeInferencerD   c                 C   s  i d| j d| jd| jd| jd| jd| jd| jd| jd	| jd
| j	d| j
d| jd| jd| j d| jd| jd| j i d| j d| jd| jd| jd| jd| jd| jd| jd| jd| jd| j d| j d| j d| jd | jd!| jd"| ji d#| jd$| jd%| jd&| jd'| jd(| jd)| jd*| jd+| j d,| j!d-| j"d.| j#d/| j$d0| j%d1| j&d2| j&d3| j&i d4| j'd5| j(d6| j)d7| j d8| j*d9| j+d:| j,d;| j-d<| j d=| j.d>| j d?| j/d@| j0dA| j1dB| j2dC| j3dD| j4| j5| j6| j7dE| _8| j| j9| j:| j;| j<| j<| j=| j>| j?| j;| j;| j,dF| _@dG| _Ai | _Bi | _Ci | _D|| _E|| _F|| _G|| _HdH| _I|| _Jd S )INAddArrayFeatureExtractorAveragePoolBatchNormalizationCastCategoryMapperCompressConcatConcatFromSequenceConstantConstantOfShapeConvCumSumDivEinsumExpandEqualFloorGatherGatherElementsGatherNDIdentityIfLoopMatMulMatMulInteger16MaxPoolMaxMinMulNonMaxSuppressionNonZeroOneHotPadRange
Reciprocal	ReduceSum
ReduceProdReshapeResizeRoundScanScatterElements
SequenceAtSequenceInsertShapeSizeSliceSoftmaxCrossEntropyLossSoftmaxCrossEntropyLossInternal!NegativeLogLikelihoodLossInternalSplitSplitToSequenceSqueezeSubTileTopK	Transpose	UnsqueezeWhereZipMapNeg	AttentionBiasGeluEmbedLayerNormalizationFastGeluGeluLayerNormalization)LongformerAttentionPythonOpSkipLayerNormalization)	embedding
bitwise_ordiagonalmax_pool2d_with_indicesmaxminmultinomialunfoldargmax
avg_pool2d_adaptive_avg_pool2dnumpy_TTr   )K_infer_symbolic_compute_ops_infer_ArrayFeatureExtractor_infer_Pool_infer_BatchNormalization_infer_Cast_infer_CategoryMapper_infer_Compress_infer_Concat_infer_ConcatFromSequence_infer_Constant_infer_ConstantOfShape_infer_Conv_pass_on_shape_and_type_infer_Einsum_infer_Expand_infer_Gather_infer_GatherElements_infer_GatherND	_infer_If_infer_Loop_infer_MatMul_infer_MatMulInteger_infer_NonMaxSuppression_infer_NonZero_infer_OneHot
_infer_Pad_infer_Range_infer_ReduceSum_infer_ReduceProd_infer_Reshape_infer_Resize_infer_Scan_infer_ScatterElements_infer_SequenceAt_infer_SequenceInsert_infer_Shape_infer_Size_infer_Slice_infer_SoftmaxCrossEntropyLoss_infer_Split_infer_SplitToSequence_infer_Squeeze_infer_Tile_infer_TopK_infer_Transpose_infer_Unsqueeze_infer_ZipMap_infer_Attention_infer_BiasGelu_infer_EmbedLayerNormalization_infer_FastGelu_infer_Gelu_infer_LayerNormalization_infer_LongformerAttention_infer_PythonOp_infer_SkipLayerNormalizationdispatcher__infer_aten_bitwise_or_infer_aten_diagonal_infer_aten_pool2d_infer_aten_minmax_infer_aten_multinomial_infer_aten_unfold_infer_aten_argmaxaten_op_dispatcher_run_suggested_merge_symbolic_dims_input_symbols_auto_merge_guess_output_rank_verbose_int_max_subgraph_id_prefix_)selfint_max
auto_mergeguess_output_rankverboseprefixr   r   r   __init__w   sH  	
 !"#$%&'()*+,-./0123456789:;<=>@ABCDEF
K
zSymbolicShapeInference.__init__Fc           	         s  t  fdd|D sJ t|} j D ]\}}||v r(|| || qd }|D ]
}t|r7|} nq-|d u rJ|D ]}| jv rI|} nq>|d u ra|D ]}t j	| t
jkr`|} nqP|d u r jdkrutdd| t|}dd |D }||t| }|| |D ]9}||krqt|rt|rt|t|ksJ t|rt|n| j|<  j D ]\}}||kr| j|< qq|rՈ jr׈   d S d S d S )Nc                    s*   g | ]}t |tkr| jv pt|qS r   )r   r   r   r1   r	   sr   r   r   r      s   * z?SymbolicShapeInference._add_suggested_merge.<locals>.<listcomp>r   z9Potential unsafe merge between symbolic expressions: ({}),c                 S   r$   r   rM   r   r   r   r   r      r'   )allsetr   itemsremoveaddr1   r   r   r   r=   Symbolr   loggerwarningformatjoinrE   indexr   r2   r   _apply_suggested_merge)	r   symbolsapplykrU   map_tor   symbols_listlensr   r   r   _add_suggested_merge   s\   






z+SymbolicShapeInference._add_suggested_mergec                 C   s|   | j sd S t| jjj|rg nt| jjj D ]$}|jjjj	D ]}|j
| j v r:| j |j
 }t|r7t||_q||_
qqd S r-   )r   rE   out_mp_graphinput
value_infor   r   r#   r   	dim_paramr1   r2   	dim_value)r   graph_input_onlyr4   r&   rU   r   r   r   r    s   (z-SymbolicShapeInference._apply_suggested_mergec                 C   s   t  | _| j| tdd t| jjjD | _tdd | jjj	D | _
tdd t| jjjD | _| jtdd | jjj	D  d S )Nc                 S      g | ]}|j |fqS r   r   r3   r   r   r   r         z6SymbolicShapeInference._preprocess.<locals>.<listcomp>c                 S   r  r   r   r3   r   r   r   r     r  c                 S   r  r   r   r3   r   r   r   r     r  c              	   S   s*   g | ]}|j t|j |jt|jfqS r   )r   r   make_tensor_value_info	data_typerE   dimsr3   r   r   r   r     s    )r.   
ModelProtor  CopyFromdictrE   r  r  graph_inputs_initializerinitializers_	known_vi_update)r   in_mpr   r   r   _preprocess  s   
z"SymbolicShapeInference._preprocessc                    s>  t dd  D smjrktt }dd |D }t|dks!J t|dkrS|d}jdkrHtd	|d | ||d d   ||  j
|dd || S jdkrgtd		|dd  |d   d S d S t  fd
d D r| d S fdd D t fddD rd jv sJ d S d S )Nc                 S      g | ]}t |tkqS r   r   r   r%   r   r   r   r   '      z9SymbolicShapeInference._merge_symbols.<locals>.<listcomp>c                 S   r$   r   r1   r%   r   r   r   r   *  r'   rL   r   z$dim {} has been merged with value {}Fallow_broadcastz!dim {} has been mergd with dim {}c                       g | ]}| d  kqS r   r   r%   r   r   r   r   =  r-  c                    s$   g | ]}| j v r j | n|qS r   )r   r%   r   r   r   r   ?     $ c                    r1  r2  r   r%   )mergedr   r   r   @  r-  )r  r   rE   r  sumr  r   r  debugr
  _check_merged_dimsr   )r   r   unique_dimsis_intint_dimr   )r   r5  r   r   _merge_symbols&  s6   


z%SymbolicShapeInference._merge_symbolsc                 C   s   g }t |}t |}t||}t|D ]Z}||k r!||d |  nd}||k r/||d |  nd}	|dks9||	kr<|	}
n,|	dkrC|}
n%| ||	g}
|
sh| jrY| j||	gdd ntdt| d t|	  |
g| }q|S )NrL   Tr  zunsupported broadcast between  )	rM   r   ranger<  r   r  r  r	  r   )r   shape1shape2	new_shaperank1rank2new_rankr4   dim1dim2new_dimr   r   r   _broadcast_shapesG  s$   
z(SymbolicShapeInference._broadcast_shapesc                 C   sD   |j | }|| jv r| j| }t|S || jv sJ t| j| jS r-   )r  r'  r,   r&  rE   r   )r   r   idxr   r+   r   r   r   
_get_shape`  s   


z!SymbolicShapeInference._get_shapec                 C   s   t | ||S r-   )rM   rK  )r   r   rJ  r   r   r   _get_shape_ranki  s   z&SymbolicShapeInference._get_shape_rankc                 C   sh   g }|  ||D ])}t|tkr&||| jv r| j| ntj|ddd qd |ks,J || q|S )NTintegernonnegative)rK  r   r   appendr   r=   r  )r   r   rJ  r6   r&   r   r   r   _get_sympy_shapel  s   z'SymbolicShapeInference._get_sympy_shapec                 C   sF   |j | }|| jv s|| jv sJ || jv r| j| S t| j| S r-   )r  sympy_data_r&  r   to_arrayr   r   rJ  r   r   r   r   
_get_valuez  s   
$z!SymbolicShapeInference._get_valuec                 C   s@   |t |jkr	d S |j| }|| jv s|| jv r| ||S d S r-   )rM   r  rR  r&  rU  rT  r   r   r   _try_get_value  s   
z%SymbolicShapeInference._try_get_valuec                 C   s~   t |D ]8\}}t|s<t|tks<t|}|| jv r.t| j| r#q| j| j|  ||< qt|| jvr<|| jt|< qd S r-   )	enumerater1   r   r   r   r   )r   new_sympy_shaper4   rH  str_dimr   r   r   _update_computed_dims  s   
z,SymbolicShapeInference._update_computed_dimsc                    s   |j dv }|sEg }t jdkr|j dv r fdd|jD }t|gd fdd|jD dd |jD |} jj	| t
 j _tt|jD ]#}|j| } jjj }|sg|	 jjj|  n||_| j|< qLd S )	N)rn   ro   r   r   r   r   r   r   r   r   r   r   r   r   	   )r   c                    s*   g | ]}| j v r| jvr j | qS r   )r&  r$  r	   r   r   r   r   r     s    zBSymbolicShapeInference._onnx_infer_single_node.<locals>.<listcomp>tmpc                    s   g | ]	}|r j | qS r   r'  r3   r   r   r   r     r   c                 S   r$   r   )r0   r3   r   r   r   r     r'   )op_typerJ   r  r  r   
make_graphoutputtmp_mp_r  r"  r   infer_shapesr?  rM   r  r  r   r'  )r   r   
skip_inferinitializers	tmp_graphi_oor+   r   r   r   _onnx_infer_single_node  s0   


z.SymbolicShapeInference._onnx_infer_single_nodeTc                    sB   j dkrtd|j|jd |j tdd t|j	t|j
 D tfdd j D tt|jdt|j
 fddD  d	d |jD }|j	fd
d jjj	D  |j	|j	  jj| t j j j j  jd t j d}|r  jd7  _d}| j  j |_|jr|  j! }|js|"  |r|#d |j
|jjj
d t$|j
  |#d |j|jjj |#d |j%|jjj% |#d |j|jjj dd |jjjD }t fdd|D }	i }
|	D ]}||j&v sJ |j&| |
|< q j&'|
 |S )N   z6Inferencing subgraph of node {} with output({}...): {}r   c                 S      g | ]}|j qS r   r   r3   r   r   r   r         z?SymbolicShapeInference._onnx_infer_subgraph.<locals>.<listcomp>c                       g | ]}| vr|qS r   r   r\  )subgraph_inputsr   r   r     r-  r]  c                    s   g | ]} j | qS r   r^  r3   r   r   r   r     r  c                 S   s   g | ]}t |jqS r   )r0   r   r3   r   r   r   r     r  c                    s   g | ]	}|j  v r|qS r   r   r3   )subgraph_implicit_inputr   r   r     r   _)r   rL   Fr  ra  r  r   c                 S   r$   r   )r,   r	   rh  r   r   r   r     r'   c                    s4   g | ]}|r|D ]}t |tkr| jvr|qqS r   )r   r   r   )r	   r   r&   r   r   r   r     s   4 )(r   r  r7  r
  r   ra  r_  r  rE   r%  r  r'  keysr   r`  r   extendr  r  rb  r"  rW   r   r   r   r   r   r   r*  r   copyr   _infer_implrR  _update_output_from_vi
ClearFieldrM   r  r   r(  )r   r   subgraphuse_node_inputinc_subgraph_idrf  symbolic_shape_inferenceall_shapes_inferredsubgraph_shapessubgraph_new_symbolic_dimsnew_dimsr&   r   )r   ro  rn  r   _onnx_infer_subgraph  sd   
" 
 


z+SymbolicShapeInference._onnx_infer_subgraphc           	         s2   fddt t jD }tdd |D rUt|D ]8\}}t|tjkr(qt|jdkr2d }nt|jdkr@t	|
 }nt|jdksIJ dd |D }|||< qdd |D }t|}|dkr|rt|D ],\}}|d u rsqjt|tkrt||k r|| ||< qjt||ksJ qj|g| ||< qj|S )Nc                       g | ]}  |qS r   )rV  r3   r   r   r   r   r     r-  z:SymbolicShapeInference._get_int_values.<locals>.<listcomp>c                 S      g | ]}|d uqS r-   r   r	   rU   r   r   r   r     r'   rL   r   c                 S   r$   r   r2   )r	   vvr   r   r   r     r'   c                 S   s$   g | ]}t |tkrt|nd qS r2  )r   rE   rM   r  r   r   r   r     r4  )r?  rM   r  r  rW  r   r:   rN   r#   r2   rO   r   rE   )	r   r   	broadcastvaluesr4   rU   new_v
values_lenmax_lenr   r  r   _get_int_values  s0   
z&SymbolicShapeInference._get_int_valuesc                    s   t |jdks	J | j|dd}tdd |D rEdd |D }t|}|r9 fddt| D | j|jd < d S  || j|jd < d S d S )	NrL   T)r  c                 S   r  r-   r   r  r   r   r   r   (  r'   zASymbolicShapeInference._compute_on_sympy_data.<locals>.<listcomp>c                 S   r+  r   )r   rE   r  r   r   r   r   )  r-  c                    s   g | ]} |qS r   r   )r	   vsop_funcr   r   r   ,  r'   r   )rM   ra  r  r  anyziprR  )r   r   r  r  is_listrT   r   r  r   _compute_on_sympy_data%  s   &z-SymbolicShapeInference._compute_on_sympy_datac                 C   s0   t |jdks|jdv sJ | |dd  d S )NrL   )r~   r   r   c                 S   s   | d S r@   r   rP   r   r   r   <lambda>6  s    z<SymbolicShapeInference._pass_on_sympy_data.<locals>.<lambda>)rM   r  r_  r  r   r   r   r   r   _pass_on_sympy_data0  s   z*SymbolicShapeInference._pass_on_sympy_datac              
   C   sH   | j |jd  }|t|jd | j |jd  jjj| 	|d d S r@   )
r'  ra  r"  r   r  r  r   r   r*   rK  )r   r   r+   r   r   r   r   8  s   
z.SymbolicShapeInference._pass_on_shape_and_typec                 C   s`   d ||}|| jv r!| j| }t|rtt|}|S |}|S tj|ddd}|| j|< |S )Nz{}_d{}TrM  )r
  r   r1   r=   r>   r2   r  r   )r   r   r   rH  rU   new_symbolic_dimr   r   r   _new_symbolic_dimB  s   


z(SymbolicShapeInference._new_symbolic_dimr   c              	   C   s,   |  d|j| jt| jjj|||S )Nz{}{}_{}_o{}_)	r  r
  r_  r   rE   r  r  r   r  )r   r   out_idxr   r   r   r   _new_symbolic_dim_from_outputL  s   z4SymbolicShapeInference._new_symbolic_dim_from_outputc                    s    fddt |D S )Nc                    s   g | ]	}  |qS r   r  r3   r   r  r   r   r   r   X  r   z>SymbolicShapeInference._new_symbolic_shape.<locals>.<listcomp>)r?  )r   rB   r   r  r   r  r   _new_symbolic_shapeW  s   z*SymbolicShapeInference._new_symbolic_shapec                 C   s  |  |d}t|jdkr'|  |d}t|d }|| d  }|d |d< nd }t|d}t|}t||d ks<J dd || d  D }t|syt| j|jd  }t|dkryt|t|ksfJ dd || d  D || d < |S t|ddg| }t|d	dg| }	d
d t||D }
t|d}|d u rdgd|  }t|dd	d}|dkr|dkrzdd t|| d  |	D }dd t|
|	|D }W n< t
y   dd t|
|	D }Y n*w |dkrg }n"dg| }nt|d| ksJ dd t|d | ||d  D }t|dd}t|D ];}|| |  }t|dkr/|||  }|r@t||
|  |	|  }n
||
|  |	|  }|d || | < q|S )Nr   rL   rj  kernel_shapec                 S   s   g | ]}t | qS r   r.  r3   r   r   r   r   i  r  zCSymbolicShapeInference._compute_conv_pool_shape.<locals>.<listcomp>c                 S      g | ]}t |qS r   r=   r>   r%   r   r   r   r   o  r  	dilationsstridesc                 S   s    g | ]\}}|d  | d  qS rL   r   )r	   r  r&   r   r   r   r   t       padsauto_pads   NOTSETutf-8VALIDNOTSETc                 S   s   g | ]
\}}t ||qS r   )r=   Modr	   r&   r   r   r   r   r   {      c                 S   s0   g | ]\}}}t d |d kr|| n|| qS r2  r   )r	   r  r   rr   r   r   r   |  s    c                 S   s   g | ]\}}t d || qS r2  r  )r	   r  r   r   r   r   r         c                 S   s   g | ]\}}|| qS r   r   )r	   p1p2r   r   r   r     r-  	ceil_mode)rQ  rM   r  r   r  r,   r'  ra  r  decode	TypeErrorr?  r=   ceiling)r   r   r6   W_shaperB   r  is_symbolic_dimsr#   r  r  effective_kernel_shaper  r  residual
total_padsr  r4   effective_input_sizestrided_kernel_positionsr   r   r   _compute_conv_pool_shapeZ  sh   
"



$z/SymbolicShapeInference._compute_conv_pool_shapec                    s>   |r	dd  D  t  fdd D s| j dd d S d S )Nc                 S   s$   g | ]}t |rt|d ks|qS r  )r1   r2   r%   r   r   r   r     r4  z=SymbolicShapeInference._check_merged_dims.<locals>.<listcomp>c                    r1  r2  r   r%   r3  r   r   r     r-  Tr=  )r  r  )r   r   r0  r   r3  r   r8    s
   z)SymbolicShapeInference._check_merged_dimsNc                 C   s6  |  |d}|  |d}t|}t|}d}d}|dkr |dks"J |dkr-|dkr-g }	n;|dkr?d}|d | |d g }	n)|dkrLd}|d | }	nd}d}| |d d |d d |d g |d g }	| j|| || gdd |d u r| j|jd  jjj}| j|j	d  }
|

t|j	d ||	 d S )Nr   rL   Fr/  )rK  rM   rI  r8  r'  r  r   r   r*   ra  r"  r   r  )r   r   output_dtype	lhs_shape	rhs_shapelhs_rankrhs_ranklhs_reduce_dimrhs_reduce_dimrB  r+   r   r   r   _compute_matmul_shape  s4   0z,SymbolicShapeInference._compute_matmul_shapec              	   C   s  t |r	|jjjn|j}t |r|jjjn|j}|j|jkrB|jr$|jn|j}td| dtjj	j
|j dtjj	j
|j |dr}tt|jj|jjD ](\}}	|	d |	d krztj }
t |sqt| ||||
_|jj| |
 qRdS || dS )zh
        update dst_tensor_type to be compatible with src_tensor_type when dimension mismatches
        z	For node z:, dst_tensor_type.elem_type != src_tensor_type.elem_type: z vs r#   r   rL   N)r"   r   r*   r   r   r_  
ValueErrorr.   onnx_pbTensorProtoDataTypeNamer(   rW  r  r#   r   TensorShapeProto	Dimensionr   r  r  r"  )r   r   r  dst_typesrc_typedst_tensor_typesrc_tensor_typenode_iddidsrH  r   r   r   _fuse_tensor_type  s.   

	z(SymbolicShapeInference._fuse_tensor_typec              	   C   sd   |  |d}|  |d}| j|jd  }|t|jd | j|jd  jjj	|d d |  d S Nr   rL   r  
rK  r'  ra  r"  r   r  r  r   r   r*   )r   r   
data_shapeindices_shaper+   r   r   r   r     s   z3SymbolicShapeInference._infer_ArrayFeatureExtractorc                    sn   dd dd dd dd  fdd fdddd d	d d
d dd d
}|j |v s,J  |||j   d S )Nc                 S   s   | d | d  S Nr   rL   r   lr   r   r   r        zDSymbolicShapeInference._infer_symbolic_compute_ops.<locals>.<lambda>c                 S   s   | d | d  S r  r   r  r   r   r   r    r  c                 S   s   | d | d kS r  r   r  r   r   r   r    r  c                 S   s   t | d S r@   )r=   floorr  r   r   r   r    s    c                    sd   t | d rt| d  j k r| d S t | d r(t| d  j k r(| d S t| d | d S r  )r1   r2   r   r=   rs   r  r   r   r   r    s
   

<c                    s`   t | d rt| d  jkr| d S t | d r&t| d  jkr&| d S t| d | d S r  )r1   r2   r   r=   rt   r  r   r   r   r    s
   

:c                 S   s   | d | d  S r  r   r  r   r   r   r    r  c                 S   s   | d | d  S r  r   r  r   r   r   r    r  c                 S   s   | d r| d S | d S )Nr   rL   rj  r   r  r   r   r   r    r-  c                 S   s
   | d  S r@   r   r  r   r   r   r    s   
 )
rX   re   rh   ri   rs   rt   ru   r   r   r   )r_  r  )r   r   funcsr   r   r   r     s   

z2SymbolicShapeInference._infer_symbolic_compute_opsc                 C      |  | d S r-   )r  r  r   r   r   r     r8   z"SymbolicShapeInference._infer_Castc              
   C   sj   | j |jd  jjj}|tjjkrtjj}ntjj}| j |j	d  }|
t|j	d || |d d S r@   )r'  r  r   r   r*   r.   r  STRINGINT64ra  r"  r   r  rK  )r   r   
input_typeoutput_typer+   r   r   r   r     s   
&z,SymbolicShapeInference._infer_CategoryMapperc                 C   s   |  |d}t| |}t|d}|d kr|g}n|}||t|t|< | j|jd  }|t	
|jd | j|jd  jjj| d S )Nr   rA   )rK  r   r  r   rC   rM   r'  ra  r"  r   r  r  r   r   r*   )r   r   input_shapecompress_lenrA   output_shaper+   r   r   r   r     s   
z&SymbolicShapeInference._infer_Compressc                    s  t fddjD rV}tdd |D rVdtdks#J g jjd < ttjD ]#}|| }t	|t
krJjjd  | q2jjd  | q2d}ttdt|}tdtjD ]}|}|r|| ||  ||< qn| tt|D ]>  |krq fddttjD tfddD rq}	t	|	tkr|	rÈj|	 nd | < q|	| < qjjd  }
|
tjd jjd  j	jjt| d S )	Nc                    s    g | ]}| j v p| jv qS r   )rR  r&  r3   r   r   r   r     r  z8SymbolicShapeInference._infer_Concat.<locals>.<listcomp>c                 S   r  r-   r   r  r   r   r   r     r'   r   rA   rL   c                    s(   g | ]} |r |  qS r   rK  )r	   i_idx)r&   r   r   r   r   r   4  s   ( c                    r1  r2  r   r%   r3  r   r   r   5  r-  )r  r  r  r  r   rR  ra  r?  rM   r   rE   rs  rP  rQ  rC   rZ  r<  r   r   r'  r"  r   r  r   r*   r7   )r   r   r  r4   r   r6   rA   r  r  r5  r+   r   )r&   r   r   r   r   r     sH   

 

z$SymbolicShapeInference._infer_Concatc                 C   s   |  |d}t|drdnd}tt|dt|| }t| |d|}|}|r8|d | |g ||d   }n|||< | j|jd  }|t	
|jd | j|jd  jjjjj| d S )Nr   new_axisrL   rA   )rK  r   rC   rM   r   r  r'  ra  r"  r   r  r  r   r   r*   r   )r   r   	seq_shaper  rA   
concat_dimrB  r+   r   r   r   r   E  s     z0SymbolicShapeInference._infer_ConcatFromSequencec                 C   s$   t |d}t|| j|jd < d S )Nr   r   )r   r   rS  rR  ra  )r   r   tr   r   r   r   X  s   
z&SymbolicShapeInference._infer_Constantc                 C   s   |  |d }| j|jd  }|d urPt|tkr|g}| | |jjjtj	j
krOtdd |D rOtjdd |D tjdtt|dd | j|jd < n| | |dd |}|t|jd |jjjt| d S )Nr   c                 S   r$   r   r.  r	   rQ   r   r   r   r   d  r'   zASymbolicShapeInference._infer_ConstantOfShape.<locals>.<listcomp>c                 S   r$   r   r  r  r   r   r   r   f  r'   )dtyper   )r  r'  ra  r   rE   rZ  r   r*   r.   r  r  r  r:   onesr;   r   rS  r   rR  r  rK  r"  r   r  r7   r   r   r6   r+   r   r   r   r   \  s*   
$z-SymbolicShapeInference._infer_ConstantOfShapec                 C   sL   |  |}| | | j|jd  }|t|jd |jjj	t
| d S r@   )r  rZ  r'  ra  r"  r   r  r   r   r*   r7   r  r   r   r   r   u  s   

z"SymbolicShapeInference._infer_Convc                 C   sJ  t |d}|dd}|d}|dkr|d | n|}d}d}d}i }|d}	|	D ]W}
|
d}| ||}t|}|dkrP|dkrL|t|
 d	 }|d
 }td
|d
 D ]&}|
|  }|dkr}||  }|| vrr|||< qWt|t	j
kr}|||< qW|d
 }q+g }ddlm} | }|dkr||d d  }|d}|dkrt|D ]	}|||  q|D ]}|dkr|||  qnAt|D ]	}|||  q|D ]}|dkr|dkr||v r|| d
 ||< qd
||< q| D ]\}}|d
kr|||  q| j|jd  jjj}| j|jd  }|t|jd || d S )Nequation        s   ->r  r      ,s   ...   rL   .   )OrderedDictrj  ,   )r   replacefindsplitrK  rM   r?  rr  r   r=   r  collectionsr  rP  r  r'  r  r   r*   ra  r"  r   r  )r   r   r  	mid_indexleft_equationnum_operandsnum_ellipsisnum_ellipsis_indicesletter_to_dimtermstermellipsis_indexr#   rB   r4   letterr   rX  r  num_letter_occurrencesright_equationright_ellipsis_indexckeyr   r  r+   r   r   r   r     sp   









z$SymbolicShapeInference._infer_Einsumc                 C   s   t | |ddd}|d urA| | | |d}| |t|}| j|jd  }|t	
|jd | j|jd  jjj| d S d S )NrL   TrS   r   )rT   rV  rZ  rK  rI  r7   r'  ra  r"  r   r  r  r   r   r*   )r   r   expand_to_shaper#   rB  r+   r   r   r   r     s   
z$SymbolicShapeInference._infer_Expandc              
      st  |  |d}tt|ddt|}|  |d}| j|jd  }|t|jd | j|j	d  j
jj|d | | ||d d    |j	d | jv rt|dkrdt|ddkr| |d}|d ur| j|j	d   t
 tkrt
|tjkrt|jdkr fdd|D | j|jd < d S  t| | j|jd < d S |dks|dksJ  | j|jd < d S d S d S d S d S )Nr   rA   rL   c                    s   g | ]} t | qS r   r  r3   datar   r   r     r-  z8SymbolicShapeInference._infer_Gather.<locals>.<listcomp>r  )rK  rC   r   rM   r'  ra  r"  r   r  r  r   r   r*   rR  rV  rE   r:   rN   r#   r2   )r   r   r  rA   r  r+   rJ  r   r  r   r     s.   ,"z$SymbolicShapeInference._infer_Gatherc                 C   sL   |  |d}| j|jd  }|t|jd | j|jd  jjj	| d S rK   r  )r   r   r  r+   r   r   r   r        z,SymbolicShapeInference._infer_GatherElementsc           	      C   s   |  |d}t|}|  |d}t|}|d }t|r ||ks"J |d d ||d   }| j|jd  }|t|jd | j|jd  j	j
j| d S r  )rK  rM   r1   r'  ra  r"  r   r  r  r   r   r*   )	r   r   r  	data_rankr  indices_ranklast_index_dimensionrB  r+   r   r   r   r     s   z&SymbolicShapeInference._infer_GatherNDc           	   	   C   s0  t |dt |dg}| |d}|d ur-t|dkr$|d |d  n	|d |d  t|D ]d\}}| j||dd}tt|jD ]P}| j	|j|  }|dkra||j|  |j| |_
n| |||j|j| j |d ur|t|dkr{dndkr|j| j
|jv r|j|j| j
 | j|j
< qDq1d S )Nthen_branchelse_branchr   rL   F)ry  )r   rV  rR   r"  rW  r  r?  rM   ra  r'  r   r  r   rR  )	r   r   	subgraphscondi_subrx  subgraph_inferi_outr+   r   r   r   r     s,    z SymbolicShapeInference._infer_Ifc                 C   sT  t |d}t|jt|jksJ t|jd }t|jD ]\}}|j}|| j|j|   ||_q| || d}td|d D ]o}|j	| }	t
|	}
t|	jrk|
rjd |
v rj|j|d  jjj|	jjj d}qB|j|d  }t
|}tt||
D ]3\}}|d |d krtj }t| ||||_|jjjj| | |	jjjj| | d}q}qB|r| jdkrtd|j|j	d  | j||dd t| |}tt|j	D ]K}| j|j	|  }||j	|d   ||kr!t|jrJ |j	|d  jjjj}|jjjd	 |jjjj}|| _|t | |j	| |_qd S )
Nbodyrj  FrL   Tr   zDRerun Loop: {}({}...), because of sequence in loop carried variables)rz  r   )!r   rM   r  rW  r   r"  r'  r  r?  ra  r,   r"   r   r   r*   r  r.   r  r  r   r  r  r   r#   r   r   r  r7  r
  rw  r  rs  rE   )r   r   rx  num_loop_carriedr4   sisi_nameneed_second_inferr  soso_shapesi_shaper  r   rH  loop_iter_dimr+   subgraph_vi_dimvi_dimr   r   r   r   !  sb   


 



z"SymbolicShapeInference._infer_Loopc                 C   r  r-   )r  r  r   r   r   r   ^  r8   z$SymbolicShapeInference._infer_MatMulc                 C   s   |  |tjj d S r-   )r  r.   r  INT32r  r   r   r   r   a  s   z+SymbolicShapeInference._infer_MatMulIntegerc                 C   sD   t | |}| j|jd  }|t|jd tjj	|dg d S )Nr   r  )
r   r  r'  ra  r"  r   r  r.   r  r  )r   r   selectedr+   r   r   r   r   d  s   &z/SymbolicShapeInference._infer_NonMaxSuppressionc                 C   sV   |  |d}t| |dd}| j|jd  }|t|jd |jj	j
||g d S r  )rL  r   r  r'  ra  r"  r   r  r   r   r*   )r   r   
input_ranknz_lenr+   r   r   r   r   i  s   (z%SymbolicShapeInference._infer_NonZeroc                 C   s   |  |d}| |d}t|dd}t|t|d }t|d | t|s*| |n|g ||d   }| j|j	d  }|
t|j	d | j|jd  jjj| d S )Nr   rL   rA   r  rj  )rQ  rV  r   rC   rM   r7   r1   r  r'  ra  r"  r   r  r  r   r   r*   )r   r   r6   depthrA   rB  r+   r   r   r   r   p  s&   

z$SymbolicShapeInference._infer_OneHotc                 C   s   t | jdkrt|d}n| |d}| |d}t|}|d urDt|d| ks+J dd t||d | ||d  D }| | n| ||}| j	|j
d  jjj}| j	|jd  }|t|jd |t| d S )N
   r  rL   r   rj  c                 S   s   g | ]\}}}|| | qS r   r   )r	   r&   pad_uppad_downr   r   r   r     r  z5SymbolicShapeInference._infer_Pad.<locals>.<listcomp>)rJ   r  r   rV  rQ  rM   r  rZ  r  r'  r  r   r   r*   ra  r"  r   r  r7   )r   r   r  r6   rB   rX  	output_tpr+   r   r   r   r     s"   z!SymbolicShapeInference._infer_Padc              	   C   sR   |  |}| | |jD ]}|sq| j| }|t||jjj	t
| qd S r-   )r  rZ  ra  r'  r"  r   r  r   r   r*   r7   )r   r   r6   rh  r+   r   r   r   r     s   



z"SymbolicShapeInference._infer_Poolc                 C   sh   |  |d}|  |d}| ||}| j|jd  }| j|jd  }|t|jd |jj	j
| d S r  )rK  rI  r'  r  ra  r"  r   r  r   r   r*   )r   r   shape0r@  rB  t0r+   r   r   r   r     s   $z-SymbolicShapeInference._infer_aten_bitwise_orc                 C   s:  |  |d}t|}| |d}| |d}| |d}|d ur(|d ur(|d us*J t||}t||}g }t|D ]\}}	|||fvrI||	 q:|| }
|| }|dkrctdt|
|| }ntdt|
| |}|| |j	d r| j
|j	d  }|t|j	d | j
|jd  jjjt| d S d S Nr   rL   rj  r  )rQ  rM   rV  rC   rW  rP  r=   rs   rt   ra  r'  r"  r   r  r  r   r   r*   r7   )r   r   r6   rB   offsetrF  rG  rB  r   valr@  rA  
diag_shaper+   r   r   r   r     s:   




z+SymbolicShapeInference._infer_aten_diagonalc           	      C   s   |  |d}t|}|dv sJ | |d}|d }|r|nt| |d|}|d d |g }| j|jd  }|t	|jd t
jjt| d S )Nr   )rL   rj  rL   r  )rQ  rM   rV  r   r  r'  ra  r"  r   r  r.   r  r  r7   )	r   r   r6   rB   num_samplesr  last_dimr  r+   r   r   r   r     s   z.SymbolicShapeInference._infer_aten_multinomialc              	      s     d}t|dksJ  fdddD |dd < | t jD ]+\}}|s-q&j| }|dkr:tjjn
j j	d  j
jj}|t||t| q&d S )Nr      c                    s   g | ]	}  d |qS r2  r  r3   r  r   r   r     r   z=SymbolicShapeInference._infer_aten_pool2d.<locals>.<listcomp>rj  r  r  rL   )rQ  rM   rZ  rW  ra  r'  r.   r  r  r  r   r   r*   r"  r   r  r7   )r   r   r6   r4   rh  r+   r*   r   r  r   r     s   

&z)SymbolicShapeInference._infer_aten_pool2dc           	      C   s`  | j |jd  }t|jdkr'|t|jd | j |jd  jjj	g  d S t|jdks0J | 
|d}|d us<J | 
|d}|d u rY| |d}| |rR|n|d |}n$| |d}t|t|}|d | }|rs|dg7 }|||d d  7 }t|}|t|jd | j |jd  jjj	| | j |jd  }|t|jd tjj| d S )Nr   rL   r  rj  )r'  ra  rM   r  r"  r   r  r   r   r*   rV  rL  r  rQ  rC   r7   r.   r  r  )	r   r   r+   keepdimr   rB   r  r#   vi1r   r   r   r     s8   
"z)SymbolicShapeInference._infer_aten_minmaxc                 C   s   |  |d}| |d}| |d}| |d}|d ur>|d ur>|d ur>|t|k s,J || | | d ||< || nt|}| |d |}| | |jd rv| j|jd  }|t	
|jd | j|jd  jjjt| d S d S r3  )rQ  rV  rM   rP  r  rZ  ra  r'  r"  r   r  r  r   r   r*   r7   )r   r   r6   	dimensionsizesteprB   r+   r   r   r   r     s*   

z)SymbolicShapeInference._infer_aten_unfoldc                 C   s   d }|j d dkrg }nE| |d}| |d}|d urQ| |d}|d ur8t|t|}|r4d||< n||= nt|}| |rB|n|d |}| | t|}|jd rs|d uru| j	|jd  }|
t|jd tjj| d S d S d S )NrL   rD   rj  r   )r  rV  rQ  rC   rM   r  rZ  r7   ra  r'  r"  r   r  r.   r  r  )r   r   rB  r   r;  r6   rB   r+   r   r   r   r   *  s(   

"z)SymbolicShapeInference._infer_aten_argmaxc                 C   sD   |  | dD ]}|t|jk r|j| dkr| j |d|d qd S )N)rL   rj  r  r9  rD   rL   )input_indexoutput_index)_propagate_shape_and_typerM   ra  )r   r   r4   r   r   r   r   C  s   
z0SymbolicShapeInference._infer_BatchNormalizationc                 C   s   | j |jd  }| |}tdd |D r7t|d }t|d }t|d }tt|| | dg}n| |g}| 	| |
t|jd | j |jd  jjjt| d S )Nr   c                 S   r  r-   r   r3   r   r   r   r   O  r'   z7SymbolicShapeInference._infer_Range.<locals>.<listcomp>rL   rj  )r'  ra  r  r  rR   r=   rs   r  r  rZ  r"  r   r  r  r   r   r*   r7   )r   r   r+   
input_datastartlimitdeltarX  r   r   r   r   L  s    

z#SymbolicShapeInference._infer_Rangec                    s&  t |dd}t| jdkrt|jdkr| |d}| j|jd  }|d u rL|s*J |t	
|jd | j|jd  jjjt| | |d| d S | |d g } fdd|D }t D ]\}}||v rq|rp|d qa|| qa|t	
|jd | j|jd  jjj| d S d S d S )NkeepdimsrL      r   c                       g | ]	}t |t qS r   rC   rM   r	   ar#   r   r   r   q  r   z;SymbolicShapeInference._infer_ReduceSum.<locals>.<listcomp>)r   rJ   r  rM   r  rV  r'  ra  r"  r   r  r   r   r*   r7   r  rL  rK  rW  rP  )r   r   	keep_dimsaxesr+   r  r4   r&   r   rM  r   r   _  s<   
z'SymbolicShapeInference._infer_ReduceSumc                 C   sb   t |d}t |dd}|dkr+|dgkr-| |d }|d ur/t|| j|jd < d S d S d S d S )NrO  rG  rL   r   )r   r  rV   rR  ra  )r   r   rO  rN  r  r   r   r   r     s   
z(SymbolicShapeInference._infer_ReduceProdc                 C   s  |  |d}| j|jd  }|d u rA| |d}t|dks J |d }t|s*J |t|jd |j	j
jt| || n| |d}td}|D ]}|| }qMg }	d}
td}t|D ]7\}}t	|tjkrq|	| n|dkr|	||  |||  }n|	| |dkr|}
q`|dkr|| }q`|	ddk sJ d|	v r|| }||	|
< | |	 |t|jd |j	j
jt|	 | | d S )NrL   r   r  rj  )rV  r'  ra  rK  rM   r1   r"  r   r  r   r   r*   r7   r  rQ  r2   rW  r=   r  rP  countrZ  r  )r   r   shape_valuer+   shape_shape
shape_rankinput_sympy_shapetotalr&   rX  deferred_dim_idxnon_deferred_sizer4   rH  r   r   r   r     s\   


z%SymbolicShapeInference._infer_Reshapec                 C   s  | j |jd  }| |d}t| jdkrJ| |d}|d urHdd t||D }| | |t	
|jd | j |jd  jjjt| d S d S | |d}| |d}| |d}|d urmdd |D }| | nT|d urt|}t|d	d
krt|d| ksJ t|d | }	t||d  }
n
dg| }	dg| }
t|}dd t||	|
|D }| | n
| | |d|}|t	
|jd | j |jd  jjjt| d S )Nr   r-  rL   c                 S   s$   g | ]\}}t t || qS r   r=   simplifyr  r  r   r   r   r     r4  z8SymbolicShapeInference._infer_Resize.<locals>.<listcomp>rj  r  c                 S   s   g | ]
}t t |qS r   rX  r   r   r   r   r     r  coordinate_transformation_modetf_crop_and_resizec              	   S   s0   g | ]\}}}}t t |||  | qS r   rX  )r	   r&   rD  endscaler   r   r   r     s    
)r'  ra  rQ  rJ   r  rV  r  rZ  r"  r   r  r  r   r   r*   r7   rM   r   rE   r  rL  )r   r   r+   rT  scalesrX  roisizesrB   	roi_startroi_endr   r   r   r     sT   


z$SymbolicShapeInference._infer_Resizec                    s  t  d}t  d}t  ddg| }t j|  fddt|D }t|jt jks3J |jd t j }t|D ]*\}}|j}|j j|   |krh|jjj	j
}	|	|	||    ||_qA | t j }
t  ddg|
 }tj jd  j|d  }	t jD ]M\}}j| }|krt|j| j}t||  t|d	 }|d | |	g ||d   }|t||j| jjj| n||j|  ||_qd S )
Nr  num_scan_inputsscan_input_axesr   c              	      s&   g | ]\}}t | | qS r   )rC   rL  )r	   r4   axr   num_scan_statesr   r   r   r     s    z6SymbolicShapeInference._infer_Scan.<locals>.<listcomp>scan_output_axesr  rL   )r   rM   r  rW  r   r"  r'  r   r   r#   r   r  r  ra  r)   rC   r   r  r*   )r   r   rx  rc  rd  rn  r4   r  subgraph_namescan_input_dimnum_scan_outputsrh  rh  r+   r#   rH  r   rf  r   r     s<   


"z"SymbolicShapeInference._infer_Scanc                 C   sL   |  |d}| j|jd  }|t|jd | j|jd  jjj	| d S r@   r  )r   r   r  r+   r   r   r   r     r  z-SymbolicShapeInference._infer_ScatterElementsc                 C   s|   |  |d}| j|jd  }|d ur:t|D ]%\}}|d urqtj }t| |d||_	|j
jjj| | qd S d S r@   )rK  r'  ra  rW  r.   r  r  r   r  r  r   r   r#   r   r"  )r   r   r  r+   r  r&   rH  r   r   r   r     s   
z(SymbolicShapeInference._infer_SequenceAtc                 C   s^   | j |jd  }| j |jd  }| j |jd  }|| |jd |_| |d|j|j d S r  )r'  r  ra  r"  r   r  r   )r   r   vi_seq	vi_tensor
vi_out_seqr   r   r   r   &  s   
z,SymbolicShapeInference._infer_SequenceInsertc                 C   s   |  |d| j|jd < d S r@   )rQ  rR  ra  r  r   r   r   r   /  s   z#SymbolicShapeInference._infer_Shapec                 C   sN   |  |d}t|| j|jd < | j|jd  t|jd tj	j
g  d S r@   )rQ  rV   rR  ra  r'  r"  r   r  r.   r  r  )r   r   r6   r   r   r   r   2  s
   z"SymbolicShapeInference._infer_Sizec                    s@  dd   fdd}t jdkr3t|d}t|d}t|d}|s+ttt|}d	gt| }n`t|d	d
d}t|dd
d}|d}|d}|d u rn|d u r_|d u snttdt|d uri|n|}|d u r|d u rz|d u sd	gt|d ur|n| }t|d
d}t|d
d}|d}|d u s|d u r|d u rtt|D ]}	|d|||< qnt
|}|D ]}	|d|||< qnt||||D ]\}}	}
}||
|| }
t|
r/|
jkr|| }
ny|
j kr|	dkrdnd}
njt|| r|
dk rtd|
||  }
t|
|| }
nM|
dkr.|
d	kr,t|
|| n|
}
n8t|| r?t|
|| }
n(z |
|| sL|| }
W n tyf   td|
||  || }
Y nw ||	|| }	t|| rt|	rtdt|	|| }	t|
|	 | |dkrdnd	 | ||< q҈| j|jd  }|t|jd |jjjt
| |j d j!v rdg|krt|d	krt|d	krt|d	krj!|j d  }t|tkst|t"j#krt|j$d	kr||d |d |d  j!|jd < d S d S d S d S d S d S d S d S )Nc                 S   s   zt | |kW S  ty   Y nw zt || kW S  ty!   Y nw z	t |  | kW S  ty4   Y nw z	t | |  kW S  tyO   t ||  dk Y S w r@   )boolr  )rQ   yr   r   r   
less_equal:  s(   z7SymbolicShapeInference._infer_Slice.<locals>.less_equalc                    sZ   z d| st | r| j kr| W S ||  W S W | S  ty,   td|  Y | S w )z/normalizes a negative index to be in [0, bound)r   zCannot determine if {} < 0)r1   r   r  r  r	  r
  )r  boundrq  r   r   r   handle_negative_indexM  s   

zBSymbolicShapeInference._infer_Slice.<locals>.handle_negative_indexr[  rO  startsendsrL   Tr  rj  r  r9  r   r  z/Unable to determine if {} <= {}, treat as equal)%rJ   r  r   rE   r?  rM   rT   rV  rQ  r  r7   r  r1   r   r   r   r=   rt   	Exceptionr  r	  r
  rY  rZ  r'  ra  r"  r   r  r   r   r*   r  rR  r:   arrayr#   )r   r   rt  rO  ru  rv  stepsrX  r4   r   er  r+   input_sympy_datar   rs  r   r   9  s   







.




*z#SymbolicShapeInference._infer_Slicec                 C   s   | j |jd  }| j |jd  jjj}||jj_|jjjt	  t
|jdkrD| |d}| j |jd  }|t|j|| d S d S r  )r'  ra  r  r   r   r*   r#   r"  r.   r  rM   rK  r   r  r   )r   r   r+   r*   r  r   r   r   r     s   
z5SymbolicShapeInference._infer_SoftmaxCrossEntropyLossc           	      C   s   |  |d}tt|ddt|}t|d}|s/t|j}|| t| g| }| | ndd |D }tt|D ]8}| j	|j|  }|
||j| | j	|jd  jjjt|d | || g ||d d    || j	|j< q<d S )Nr   rA   r  c                 S   r  r   r  r   r   r   r   r     r  z>SymbolicShapeInference._infer_Split_Common.<locals>.<listcomp>rL   )rQ  rC   r   rM   ra  r=   r>   rZ  r?  r'  r"  r  r   r   r*   r7   r   )	r   r   make_value_info_funcrT  rA   r  num_outputsrg  r+   r   r   r   _infer_Split_Common  s&   

(z*SymbolicShapeInference._infer_Split_Commonc                 C      |  |tj d S r-   )r~  r   r  r  r   r   r   r        z#SymbolicShapeInference._infer_Splitc                 C   r  r-   )r~  r   make_sequence_value_infor  r   r   r   r     r  z-SymbolicShapeInference._infer_SplitToSequencec              	      s  |  |d t| j}|dk rt|d}| |dd u sJ n| |d}t|dd u s.J |d u r_dd  D }| jdkr^dd  D }t|dkr^td|j	 d	|j
 d
d|   nV fdd|D }g }tt D ]D}||vr~| |  qp | dkst | tksJ | jdkrt | tkrtd|j	 d	|j
 d
d |  d| d  qp| j|jd  }|t|jd | j|jd  jjj| | | d S )Nr   rH  rO  rL   c                 S   s   g | ]}|d kr|qS r  r   r   r   r   r   r     r-  z9SymbolicShapeInference._infer_Squeeze.<locals>.<listcomp>c                 S   s   g | ]
}t |tkr|qS r   )r   r2   r   r   r   r   r     r  z+Symbolic dimensions in input shape of op: 'z	' node: 'z'. z8Assuming the following dimensions are never equal to 1: c                    rI  r   rJ  rK  r  r   r   r     r   zAssuming the dimension 'z' at index z of the input to be equal to 1.)rK  rJ   r  r   rV  r   rM   r  r7  r_  r   r?  rP  r   r2   r'  ra  r"  r   r  r  r   r*   r  )r   r   op_setrO  r  symbolic_dimensionsr4   r+   r   r  r   r     sP   


 z%SymbolicShapeInference._infer_Squeezec           	      C   s   |  |d}g }|d ur,| |d}t|D ]\}}|||  }|| q| | n
| | |d|}| j|jd  }|	t
|jd |jjjt| d S rK   )rV  rQ  rW  rP  rZ  r  rL  r'  ra  r"  r   r  r   r   r*   r7   )	r   r   repeats_valuerX  rT  r4   r&   rH  r+   r   r   r   r   	  s"   z"SymbolicShapeInference._infer_Tilec           	      C   s   |  |d}tt|dd|}| |d}t| jdkr"t|d}n| |d }|d kr3| |}nt|}t	|t
tfv rD|||< n| |d}|||< | | t|}tt|jD ]}| j|j|  }|t|j| |j	jj| q^d S )Nr   rA   r  r[  r  rL   )rL  rC   r   rK  rJ   r  r  r  rR   r   r2   r   rQ  rZ  r7   r?  rM   ra  r'  r"  r   r  r   r*   )	r   r   rB   rA   rB  r  rX  rg  r+   r   r   r   r     s*   
"z"SymbolicShapeInference._infer_TopKc                 C   s   |j d | jv r?| |d}t|dtttt|}| j|j d  }tj	t
|j| t|d  | j|jd < d S d S )Nr   perm)rO  )r  rR  rK  r   reversedrE   r?  rM   r:   	transposerx  reshapetupleflattentolistra  )r   r   r  r  rC  r   r   r   r   :  s   $z'SymbolicShapeInference._infer_Transposec           	         s  |  |d}t| j}|dk rt|d}| |dd u sJ n| |d}t|dd u s.J t|t|   fdd|D }d}g }t D ]}||v rS|d qG|||  |d7 }qG| j|j	d  }|
t|j	d | j|jd  jjj| | | d S )Nr   rH  rO  rL   c                    s   g | ]}t | qS r   )rC   rK  output_rankr   r   r   P  r  z;SymbolicShapeInference._infer_Unsqueeze.<locals>.<listcomp>)rK  rJ   r  r   rV  rM   r?  rP  r'  ra  r"  r   r  r  r   r   r*   r  )	r   r   r  r  rO  
input_axisr  r4   r+   r   r  r   r   C  s2   


z'SymbolicShapeInference._infer_Unsqueezec                 C   s   d }t |dd urtjj}nt |dd urtjj}|d usJ t }|jd |_tjj|j	j
jjjj_||j	j
jj_| j|jd  }|| d S )Nclasslabels_int64sclasslabels_stringsr   )r   r.   r  r  r  r/   ra  r   FLOATr   r   r*   map_type
value_typer   key_typer'  r"  )r   r   map_key_typenew_vir+   r   r   r   r   f  s   
z$SymbolicShapeInference._infer_ZipMapc           
      C   s  |  |d}|  |d}t|dkrt|dksJ t|d}|d ur4t|dks+J t|d |d< n
t|d d |d< | j|jd  jjj}| j|j	d  }|
t|j	d || t|j	dkr|  |d}|  |d}|  |d}	t|dkrt|	dv r|	d	 |d< n&t|d trt|d tr|d |d  |d< n|d  d
|d  |d< | j|j	d  }|
t|j|| d S d S d S )Nr   rj  r  rL   qkv_hidden_sizesr9     r:  r  +)rK  rM   r   r2   r'  r  r   r   r*   ra  r"  r   r  
isinstancer   )
r   r   r#   
shape_biasqkv_hidden_sizes_attrr  r+   r  
past_shape
mask_shaper   r   r   r   u  s2   
z'SymbolicShapeInference._infer_Attentionc                 C   r  r-   rB  r  r   r   r   r     r8   z&SymbolicShapeInference._infer_BiasGeluc                 C   r  r-   r  r  r   r   r   r     r8   z&SymbolicShapeInference._infer_FastGeluc                 C   r  r-   r  r  r   r   r   r     r8   z"SymbolicShapeInference._infer_Geluc                 C   r  r-   r  r  r   r   r   r     r8   z0SymbolicShapeInference._infer_LayerNormalizationc                 C   r  r-   r  r  r   r   r   r     r8   z1SymbolicShapeInference._infer_LongformerAttentionc                 C   s   |  |d}|  |d}t|dkrt|dksJ ||d g }| j|jd  jjj}| j|jd  }|t	
|jd || |d g}| j|jd  }|t	
|jd tjj| t|jdkr{| j|jd  }|t	
|jd || d S d S )Nr   rj  rL   )rK  rM   r'  r  r   r   r*   ra  r"  r   r  r.   r  r(  )r   r   input_ids_shapeword_embedding_shaper  word_embedding_dtyper+   mask_index_shaper   r   r   r     s   
z5SymbolicShapeInference._infer_EmbedLayerNormalizationc                 C   r  r-   r  r  r   r   r   r     r8   z4SymbolicShapeInference._infer_SkipLayerNormalizationc           	      C   s   t |d}|s	J t |d}|sJ | j|jd  }|t|jd tjjg  t	t
|jd D ]+}| j|j|d   }| || |}t|}t|j|d  || |}|| q2d S )Noutput_tensor_typesoutput_tensor_ranksr   rL   )r   r'  ra  r"  r   r  r.   r  r  r?  rM   r  r7   )	r   r   r  r  r+   r4   r6   r#   r  r   r   r   r     s   

z&SymbolicShapeInference._infer_PythonOpc                 C   sP   |  ||}| j|j|  jjj}| j|j|  }|t	|j| || d S r-   )
rK  r'  r  r   r   r*   ra  r"  r   r  )r   r   r@  rA  r#   r  r+   r   r   r   rB    s   z0SymbolicShapeInference._propagate_shape_and_typec                 C   s2   t |tkrdS d|vrdS || j v rdS dS )NFunk__T)r   r   r   rr  )r   r  r   r   r   _is_none_dim  s   z#SymbolicShapeInference._is_none_dimc                 C   s    |D ]}|  |r|  S qd S r-   )r  )r   	out_shapeoutr   r   r   _is_shape_contains_none_dim  s
   
z2SymbolicShapeInference._is_shape_contains_none_dimc                    s	  |pi _ jjd jdd t _jjjD ]C}t|}|d u r&qt	|j
r4|j
jjjjj}n|j
jjj}t|D ]\}}|d u rRt|j||| _q>jdd |D  qjD ]'}|jv r~j| }|jv suJ j| j|< qbtj|dddj|< qbt _jj jjd i }	fdd	jjjD ]}
|
|	|
jd
 < qg }tdd t jjjt jjj! D t"fddjjjD rjjj}ngt#fddjjjD sHt$|}jjjD ]&jd
 vrt#fdd|	jd
  D rj |% q|t$|kr:t#fddjjjD s:t&dt#fddjjjD r|D ]t#fddjD s\J ' d}j(j)v rsj)j(  n[j(dv rj*jd
  }t$|j
jjjd
krtj+j,|j
j_n:j(dkrΈj-dkrΈj.D ]*}|jdkrt/|j0t1r|j02dn|j0}|j3v rd}j3|   nqj4dkrt56j(d j  tjD ]\}}t56d7|||j8v rdnd qj(dv rJj*jd
  }t$t9|j
fddt:t$jD }t:j(dv r*dnd
 D ]  fdd|D }t$|d krHj;|dd! q.t:t$jD ]w}j*j|  }|j
}|<d"}|d#vrj4dkr|d$kr|jj<d"}d%|krt56d&7j| tt|tj+j=>|j
jjjj nt56d'7j| | nt56d(7j| | qQt||jjtj+j,k}j4dkrt56d)7j| ttj+j=>|j
jj j| j v rt56d*tj j|    d urd v s?s|rȈj@r܈j(d+v rfd,dt:t$jD }j(dv rd v s5?rd v r@Ad nA?fd-d|D }t$|d
 dkrh|d
 t$|d
 d k sjJ t$|d  dkr|d  t$|d  d k sJ nj(d.krBd
Cd g}ng }|rt:t$D ]2 d urD sqfd/d|D }t$|d
kr҈Ed0d tF||D  qd_Gnd_Gnd_GjGdkrqj(j)vrq|sq|od u pt$d
k}|rjHr
Id
nd1nt$d
krqJ|}|r,j*jd
  j
jj}n|j
jj}|tKL|j|tM| j4d
krl|rTt56d27j(j|j j4dkrlt56d)7j| t||j
jj d_GqQj4d
ks~j@r~|rt56d3j( d j  t56d4 jD ]}t56j*|  qt56d5 jD ]}t56j*|  qj@r|st56d6tj    dS qQqJd_GdS )7Nr  T)r  c                 S   s   g | ]
}t |tkr|qS r   r,  r%   r   r   r   r     r  z6SymbolicShapeInference._infer_impl.<locals>.<listcomp>)rN  positiver%  c                    s   t dd | jD }g }d| jkrt| dt| dg}n| jdv r't| dg}|D ]C}dd	 |jD  t  }|jD ]} |j q9|jD ]}| fd
d|D  qE|| |jD ]}|j|v rk|	|j q^q)|S )Nc                 s   s    | ]}|r|V  qd S r-   r   r3   r   r   r   	<genexpr>  s    zISymbolicShapeInference._infer_impl.<locals>.get_prereq.<locals>.<genexpr>rn   r  r  )ro   r   r  c                 S   s   h | ]}|j qS r   r   r3   r   r   r   	<setcomp>  rl  zISymbolicShapeInference._infer_impl.<locals>.get_prereq.<locals>.<setcomp>c                    rm  r   r   r3   g_outputs_and_initializersr   r   r     r-  zJSymbolicShapeInference._infer_impl.<locals>.get_prereq.<locals>.<listcomp>)
r  r  r_  r   r%  r   r(  ra  r   r  )r   namesr  gg_prereqnr4   )
get_prereqr  r   r    s,   






z6SymbolicShapeInference._infer_impl.<locals>.get_prereqr   c                 S   rk  r   r   r3   r   r   r   r   +  rl  c                       g | ]}|j  v qS r   r   rq  sorted_known_vir   r   r   ,  r  c                    r  r   r   rq  r  r   r   r   0  r  c                    s   g | ]}|r| v qS r   r   r3   r  r   r   r   4  r-  c                    r  r   r   rq  r  r   r   r   9  r  zInvalid model with cyclic graphc                    s   g | ]	}|r| j v qS r   r^  r3   r   r   r   r   >  r   F)ConvTransposeATenzorg.pytorch.atenoperatorr  rj  z: z  Input {}: {} {}rD   )	rX   r   ru   re   rp   MatMulIntegerrq   r   Sumc                    r  r   r  r3   r  r   r   r   j  r-  )rp   r  rq   c                    s0   g | ]}t |  kr|t |    qS r   r  r   )r&   out_rankr   r   r   l  r5   rL   r/  r   )r   sparse_tensor_typeNr   r   z  {}: sequence of {} {}z  {}: sequence of {}z  {}: {}z  {}: {} {}z  Sympy Data: )rX   r   ru   re   rp   r  rq   r_   r   r  rh   LessGreaterLessOrEqualGreaterOrEqualrt   rs   c                    r  r   r  r3   r  r   r   r     r-  c                        g | ]}t |t    qS r   r  r   rJ  r  r   r   r     r  rg   c                    r  r   r  r   r  r   r   r     r  c                 S   s8   g | ]\}}|d krt || r|| nt|| qS r2  )r1   r   )r	   r   r4   r   r   r   r     s    r  z3Possible unknown op: {} node: {}, guessing {} shapez*Stopping at incomplete shape inference at znode inputs:znode outputs:z	Merging: )NrR  r  r  rw  r  r  r   r  r,   r"   r   r   r*   r   r#   r   rW  r   r  r   r  r(  r   r   r=   r  r.   r!  rb  r"  r   ra  rE   r%  r  r  rM   rP  rw  ri  r_  r   r'  r  	UNDEFINEDrG   r   r  r   bytesr  r   r   r  r7  r
  r&  r)   r?  r8  r   r  r  r  r   r  rK  rU  r  r  r  r   r   rL  r  r   r  r7   )r   start_sympy_datar4   r  
input_dimsi_dimr   r   s_mergeprereq_for_noder  sorted_nodesold_sorted_nodes_lenknown_aten_opr+   r
   aten_op_namer   	in_shapesin_dimsrg  out_typeout_type_kindseq_cls_typeout_type_undefinedshapesdim_idxis_unknown_oprB  	out_dtyperh  r   )r&   r  rJ  r   r  r  r   r  r   ru    s  





*



  






00

 





 'z"SymbolicShapeInference._infer_implc                 C   s2   | j jjD ]}|j| jv r|| j|j  qd S r-   )r  r  ra  r   r'  r"  )r   ra  r   r   r   rv  	  s
   z-SymbolicShapeInference._update_output_from_vic                 C   sl   t | }|r
|dk rtd d S t||||}d}||  |jr)| }|js"|  |s3td|j	S )N   z.Only support models of onnx opset 7 and above.Fz#Incomplete symbolic shape inference)
rJ   r  r	  rW   r*  r   ru  rv  rw  r  )r)  r   r   r   r   
onnx_opsetr{  r|  r   r   r   rc  	  s   

z#SymbolicShapeInference.infer_shapes)rD   )F)TT)r   r   r2  )Tr-   )r  FFr   )b__name__
__module____qualname__r   r  r  r*  r<  rI  rK  rL  rQ  rU  rV  rZ  ri  r  r  r  r  r   r  r  r  r  r8  r  r  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r~  r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rB  r  r  ru  rv  staticmethodrc  r   r   r   r   rW   v   s    

d
-!	
7
;




@
	)>="!	!40#	y0	# 
	
  6rW   c                  C   s   t  } | jdddd | jddd | jdd	d
dd | jddtdd | jddd
dd | jddtdd | jddd
dd | jddd
dd | jdddd | jddtd d |  S )!Nz--inputTzThe input model file)requiredhelpz--outputzThe output model file)r  z--auto_mergez:Automatically merge symbolic dims when confliction happens
store_trueF)r  actiondefaultz	--int_maxzGmaximum value for integer to be treated as boundless for ops like slicer  )r  r   r  z--guess_output_rankz;guess output rank to be the same as input 0 for unknown opsz	--verbosezHPrints detailed logs of inference, 0: turn off, 1: warnings, 3: detailedr   z--save_as_external_dataz%Saving an ONNX model to external dataz--all_tensors_to_one_filez(Saving all the external data to one filez--external_data_locationz+The file location to save the external filez./)r  r  z--external_data_size_thresholdz$The size threshold for external datai   )argparseArgumentParseradd_argumentr2   
parse_args)parserr   r   r   parse_arguments/	  sf   r  __main__zinput model: zoutput model z!Doing symbolic shape inference...TF)save_as_external_dataall_tensors_to_one_filelocationsize_thresholdconvert_attributezDone!r-   )0r  loggingnumpyr:   r.   r=   r   r   r   	packagingr   parse__version__	getLoggerr  r  r   r   r"   r)   r,   r0   r7   r1   rC   rJ   rR   rT   rV   rW   r  argsinfor  ra  rc  loadr   r   r   r   out_mpr  
save_modelr  external_data_locationexternal_data_size_thresholdsaver   r   r   r   <module>   s   




                 J6



