a
    d                     @   s  d dl Z d dlmZmZmZmZmZmZmZ d dl	Z	d dl
Z
d dlZd dlZd dlZd dlZd dlZd dlmZ d dlmZ d dlZd dlZd dlZd dlZd dl mZmZmZmZmZmZmZ d dlmZm Z  eeefZ!e"e#Z$G dd de%Z&z8d dl'Z'd d	l(m)Z) d d
l*m+Z+ d dl,m-Z-m.Z. dZ/W n e0yD   dZ/Y n0 e j1j2j3Z3g dZ4da5G dd dZ6dd Z7dd Z8dd Z9dd Z:dd Z;dd Z<dd Z=dd  Z>d!d" Z?d#d$ Z@d%d& ZAd'd( ZBd)d* ZCd+d, ZDG d-d. d.ZEe/r6G d/d0 d0e'jFZGG d1d2 d2e'jFZHG d3d4 d4e'jFZIG d5d6 d6e'jFZJed7d8d9 ZKd:d; d<d; d=d; d>d; d?d; d@d; dAd; dBd; dCd; dD	ZLdEdF ZMdGdH ZNdIdJ ZOdKdL ZPi eLdMd; dNd; dOd; dPd; dQd; dRd; dSd; eOdTd; ePdUd; dVd; dWd; dXd; dYZQdZd[d; iZRe
jSe
jTd\ZUd]d^ ZVd_dZ ZWh d`ZXh daZYdbdchZZh ddZ[dedfhZ\dgdh Z]ej^ej_ej`ejaejbejceSeTejdejeejfejgejhdiZih djZjdbdchZkh dkZldldm Zmdndo Zndpdq ZoeQp D ]\ZqZreneqer qeRp D ]\ZqZreoeqer qdrds ZseQp D ]4\ZqZreqeYv r&eseqe neseqe eseqe q[q[rd|dtduZte/rtG dvdw dwe'juZvG dxdy dye+Zwex ZyG dzd{ d{ZzdS )}    N)SetDictListTypeOptionalcastUnion)contextmanager)	lru_cache)SymIntSymFloatSymBoolsym_not	sym_floatsym_maxsym_min)
ShapeGuardSourcec                   @   s   e Zd ZdS )GuardOnDataDependentSymNodeN)__name__
__module____qualname__ r   r   n/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/fx/experimental/symbolic_shapes.pyr      s   r   )
precedence)
StrPrinter)	fuzzy_andfuzzy_orTF)has_symbolic_sizes_stridescreate_contiguousShapeEnvSymDispatchModeFloorDiv	guard_intguard_floatguard_scalar	wrap_nodemethod_to_operatorhint_intSYMPY_INTERPc                   @   s$   e Zd Zdd Zdd Zdd ZdS )r!   c                 C   s
   t  d S N)NotImplementedError)selffunctypesargskwargsr   r   r   __sym_dispatch__F   s    z SymDispatchMode.__sym_dispatch__c                 C   s,   t }t| drt|  dn|| _| a | S )Ninnerz< has already been used as a mode. Please use a fresh version)SYM_FUNCTION_MODEhasattrRuntimeErrorr2   )r,   oldr   r   r   	__enter__I   s    
zSymDispatchMode.__enter__c                 C   s
   | j ad S r*   )r2   r3   )r,   exc_typeexc_valexc_tbr   r   r   __exit__S   s    zSymDispatchMode.__exit__N)r   r   r   r1   r7   r;   r   r   r   r   r!   E   s   
r!   c                 C   s   | j S r*   )Z_has_symbolic_sizes_strides)elemr   r   r   r   W   s    r   c                 C   s:   dg}t | d d D ]}|||d   qtt |S )N   )reversedappendlist)shapestridesdimr   r   r   r   Z   s    r   c                 C   s8   t }|sJ |ja zg }|| |||W |a S |a 0 d S r*   )r3   r2   r1   )r-   r/   r0   moder.   r   r   r   _handle_sym_dispatch`   s    rF   c                 C   s.   t | tjr| j S t| tu s*J | | S r*   )
isinstancetorchr   noderequire_hinttypeintar   r   r   r(   l   s    
r(   c                 C   sT   t | ttfrt| S t | ttfr,t| S t | ttfrBt	| S t
d|  d S )Nzunrecognized scalar )rG   r   bool
guard_boolr   rL   r#   r   floatr$   AssertionErrorrM   r   r   r   r%   r   s    r%   c                 C   s0   t | tr| jddS t| tu s,J | | S N r   )rG   r   rI   rP   rK   rO   rM   r   r   r   rP   |   s    
rP   c                 C   s0   t | tr| jddS t| tu s,J | | S rS   )rG   r   rI   r#   rK   rL   rM   r   r   r   r#      s    
r#   c                 C   s.   t | tr| jddS t | ts*J | | S rS   )rG   r   rI   r$   rQ   rM   r   r   r   r$      s    
r$   c                 C   s   t | dr|  S t| S )N__sym_sqrt__)r4   rU   mathsqrtrM   r   r   r   sym_sqrt   s    
rX   c                 C   sZ   t |tr|jS t|tu r&| |S t|tu r<| |S t|tu rR| 	|S t
S d S r*   )rG   SymTypesrI   rK   rO   	wrap_boolrL   wrap_intrQ   
wrap_floatNotImplementedr,   numr   r   r   to_node   s    



r`   c                 C   s   dd | j jD S )Nc                 S   s    g | ]}|j d kr|jd qS )placeholderval)opmeta.0nr   r   r   
<listcomp>       z'fx_placeholder_vals.<locals>.<listcomp>graphnodesgmr   r   r   fx_placeholder_vals   s    ro   c                 C   s   dd | j jD S )Nc                 S   s   g | ]}|j d kr|jqS )ra   )rc   targetre   r   r   r   rh      ri   z*fx_placeholder_targets.<locals>.<listcomp>rj   rm   r   r   r   fx_placeholder_targets   s    rq   c                 G   s   | j t| |S r*   )	shape_envevaluate_guards_for_argsro   rn   r/   r   r   r   eval_guards   s    ru   c                 G   s   | j t| |S r*   )rr   bind_symbolsro   rt   r   r   r   rv      s    rv   c                   @   s4  e Zd ZdZdceeeef  dddZe	dd Z
dd	 Ze	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dd Zd d! Zd"d# Zd d$d%d&Zd d$d'd(Zd d$d)d*Zd d$d+d,Zd d$d-d.Zd d$d/d0Zd d$d1d2Zd d$d3d4Z d d$d5d6Z!d d$d7d8Z"d d$d9d:Z#d d$d;d<Z$d d$d=d>Z%d d$d?d@Z&d d$dAdBZ'd d$dCdDZ(d d$dEdFZ)d d$dGdHZ*d d$dIdJZ+d d$dKdLZ,d d$dMdNZ-d d$dOdPZ.d d$dQdRZ/d d$dSdTZ0dUdV Z1dWdX Z2dYdZ Z3d[d\ Z4d]d^ Z5d_d` Z6dadb Z7dS )dSymNodez
    This is a type erased SymInt/SymFloat which we use to do actual operations.
    End users don't touch this.  Magic methods are NOT defined on this object.
    Nhintc                 C   sP   || _ || _|| _|d u r:| j|j| _d | _|   nd | _|| _|| _	d S r*   )
_exprrr   pytypeexprxreplace
var_to_val
_hint_expr_hint_update_hintconstant)r,   r|   rr   r{   ry   r   r   r   r   __init__   s    
zSymNode.__init__c                 C   s   |    | jS r*   )_update_exprrz   r,   r   r   r   r|      s    zSymNode.exprc                 C   s4   | j j| jj kr0| | j| j | _d | _ d S r*   )r   free_symbolsrr   replacementskeysr{   replacer   r   r   r   r   r      s    zSymNode._update_hintc                 C   s   | j d u r|   | j S r*   )r   r   r   r   r   r   ry      s    
zSymNode.hintc                 C   s>   | j d u r4|   | j d u r,| j| jq:| j S n| j S d S r*   )r   r   rr   _make_data_dependent_errorr   r   r   r   r   rJ      s    

zSymNode.require_hintc                 C   s   | j | j| _d S r*   )rr   r   rz   r   r   r   r   r      s    zSymNode._update_exprc                 C   s
   | j tu S r*   )r{   rL   r   r   r   r   is_int   s    zSymNode.is_intc                 C   s
   | j tu S r*   )r{   rQ   r   r   r   r   is_float   s    zSymNode.is_floatc                 C   s
   | j tu S r*   )r{   rO   r   r   r   r   is_bool   s    zSymNode.is_boolc                 C   s*   t |tu sJ tt|| jt||dS N)r   )rK   rL   rw   sympyIntegerrr   r^   r   r   r   r[     s    zSymNode.wrap_intc                 C   s*   t |tu sJ tt|| jt||dS r   )rK   rQ   rw   r   Floatrr   r^   r   r   r   r\     s    zSymNode.wrap_floatc                 C   s0   t |tu sJ t|rtjntj| jt||dS r   )rK   rO   rw   r   truefalserr   r^   r   r   r   rZ     s    zSymNode.wrap_boolc                 C   s   | S r*   r   r   r   r   r   clone  s    zSymNode.clonec                 C   s   | j  S r*   r|   r   r   r   r   str  s    zSymNode.strc                 C   s   |   S r*   r   r   r   r   r   __str__  s    zSymNode.__str__c                 C   s   |   S r*   r   r   r   r   r   __repr__  s    zSymNode.__repr__returnc                 C   s
   |  |S r*   )Z_addr,   otherr   r   r   add  s    zSymNode.addc                 C   s
   |  |S r*   )Z_subr   r   r   r   sub   s    zSymNode.subc                 C   s
   |  |S r*   )Z_mulr   r   r   r   mul#  s    zSymNode.mulc                 C   s
   |  |S r*   )Z_modr   r   r   r   mod&  s    zSymNode.modc                 C   s
   |  |S r*   )Z_powr   r   r   r   pow)  s    zSymNode.powc                 C   s
   |  |S r*   )Z_and_r   r   r   r   and_,  s    zSymNode.and_c                 C   s
   |  |S r*   )Z_or_r   r   r   r   or_/  s    zSymNode.or_c                 C   s
   |  |S r*   )Z_truedivr   r   r   r   truediv2  s    zSymNode.truedivc                 C   s
   |  |S r*   )Z	_floordivr   r   r   r   floordiv5  s    zSymNode.floordivc                 C   s   |   S r*   )Z_sym_notr   r   r   r   r   8  s    zSymNode.sym_notc                 C   s
   |  |S r*   )_eqr   r   r   r   eq;  s    z
SymNode.eqc                 C   s
   |  |S r*   )Z_ner   r   r   r   ne>  s    z
SymNode.nec                 C   s
   |  |S r*   )Z_gtr   r   r   r   gtA  s    z
SymNode.gtc                 C   s
   |  |S r*   )Z_ltr   r   r   r   ltD  s    z
SymNode.ltc                 C   s
   |  |S r*   )Z_ler   r   r   r   leG  s    z
SymNode.lec                 C   s
   |  |S r*   )Z_ger   r   r   r   geJ  s    z
SymNode.gec                 C   s   |   S r*   )_floorr   r   r   r   floorM  s    zSymNode.floorc                 C   s   |   S r*   )Z
_sym_floatr   r   r   r   r   P  s    zSymNode.sym_floatc                 C   s   |   S r*   )_ceilr   r   r   r   ceilS  s    zSymNode.ceilc                 C   s   |   S r*   )Z_negr   r   r   r   negV  s    zSymNode.negc                 C   s
   |  |S r*   )Z_sym_minr   r   r   r   r   Y  s    zSymNode.sym_minc                 C   s
   |  |S r*   )Z_sym_maxr   r   r   r   r   \  s    zSymNode.sym_maxc                 C   s   |   S r*   )Z	_sym_sqrtr   r   r   r   rX   _  s    zSymNode.sym_sqrtc                 G   s
   | j | S r*   )Z'_is_non_overlapping_and_dense_indicator)r,   r/   r   r   r   &is_non_overlapping_and_dense_indicatorb  s    z.SymNode.is_non_overlapping_and_dense_indicatorc                 C   s
   |  |S r*   )r   r   r   r   r   sym_orf  s    zSymNode.sym_orc                 C   s
   |  |S r*   )r   r   r   r   r   sym_andi  s    zSymNode.sym_andc                 C   s.   t | jjdkrt| jS td| j d S )Nr   z7Trying to extract a concrete int out of a symbolic int )lenr|   r   rL   r5   r   r   r   r   int_m  s    
zSymNode.int_c                 C   sF   | j | j| j}z
t|W S  ty@   td|   Y n0 d S )NzFailed to convert to int: )rr   evaluate_exprr|   ry   rL   	Exceptionlogwarningr,   filelinerr   r   r   r#   s  s    
zSymNode.guard_intc                 C   sF   | j | j| j}z
t|W S  ty@   td|   Y n0 d S )NzFailed to convert to float: )rr   r   r|   ry   rQ   r   r   r   r   r   r   r   r$   }  s    
zSymNode.guard_floatc                 C   sF   | j | j| j}z
t|W S  ty@   td|   Y n0 d S )NzFailed to convert to bool: )rr   r   r|   ry   rO   r   r   r   r   r   r   r   rP     s    
zSymNode.guard_boolc                 C   s   |  ddS rS   )rP   r   r   r   r   bool_  s    zSymNode.bool_)N)8r   r   r   __doc__r   r   rL   rQ   r   propertyr|   r   ry   rJ   r   r   r   r   r[   r\   rZ   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   rX   r   r   r   r   r#   r$   rP   r   r   r   r   r   rw      sd   !





rw   c                   @   s   e Zd Zedd ZdS )Powc                 C   s:   |j rtdS |j r.|dk r.t| dn|| S d S )Nr=   r   z% cannot be raised to a negative power)is_zeror   r   ZeroDivisionError)clsbaseexpr   r   r   eval  s
    
zPow.evalNr   r   r   classmethodr   r   r   r   r   r     s   r   c                   @   s   e Zd Zedd ZdS )TrueDivc                 C   s   |j rtdn|| S d S )Ndivision by zero)r   r   )r   r   divisorr   r   r   r     s    
zTrueDiv.evalNr   r   r   r   r   r     s   r   c                   @   sX   e Zd ZdZdZdZdZedd Zedd Z	d	d
 Z
dd Zdd Zedd ZdS )r"   z
        We maintain this so that:
        1. We can use divisibility guards to simplify FloorDiv(a, b) to a / b.
        2. Printing out the expression is nicer (compared to say, representing a//b as (a - a % b) / b)
        )   2   Tc                 C   s
   | j d S Nr   r/   r   r   r   r   r     s    zFloorDiv.basec                 C   s
   | j d S )Nr=   r   r   r   r   r   r     s    zFloorDiv.divisorc                 C   s.   | | j| j}| | j| j}| d| S )Nz//)Zparenthesizer   r   r   )r,   printerr   r   r   r   r   	_sympystr  s    zFloorDiv._sympystrc                 C   s   t | jj| jjgS r*   )r   r   is_realr   r   r   r   r   _eval_is_real  s    zFloorDiv._eval_is_realc                 C   s   t | jj| jjgS r*   )r   r   
is_integerr   r   r   r   r   _eval_is_integer  s    zFloorDiv._eval_is_integerc                    sZ   fdd}|  | j r,td j r:tjjS  jrLdkrL S  jrddkrdt S t tj	rttj	r  S t tj	tj
frttj	tj
frt  S t trt jd  jd  S t tjr" jD ]2}t|}|krt | ||    S qt }|dkrVtt | t| S d S )Nc                    sF   | j du r| jdu r| js | jrBtdt j dtj dd S )NFz%unsupported operand type(s) for //: 'z' and 'z', expected integer or real)r   r   Z
is_complexZ
is_Boolean	TypeErrorrK   r   xr   r   r   r   check_supported_type  s     z+FloorDiv.eval.<locals>.check_supported_typer   r=   r   )r   r   r   SZZeror   r   r   rG   r   r   r"   r/   Addgcdsimplify)r   r   r   r   rN   r   r   r   r   r     s6    
$


zFloorDiv.evalN)r   r   r   r   nargsr   r   r   r   r   r   r   r   r   r   r   r   r   r   r"     s   

r"   c                   @   s   e Zd ZdZedd ZdS )!IsNonOverlappingAndDenseIndicatorTc                 G   sp   t |d dksJ tdd |D rlt |d }|d| }||d  }ttdd |D dd |D S d S )Nr   r   c                 s   s   | ]}t |tjV  qd S r*   )rG   r   r   rf   rN   r   r   r   	<genexpr>  ri   z9IsNonOverlappingAndDenseIndicator.eval.<locals>.<genexpr>c                 S   s   g | ]}t |qS r   rL   rf   sr   r   r   rh     ri   z:IsNonOverlappingAndDenseIndicator.eval.<locals>.<listcomp>c                 S   s   g | ]}t |qS r   r   r   r   r   r   rh     ri   )r   allrL   !eval_is_non_overlapping_and_dense)r   r/   rD   sizesrC   r   r   r   r     s    z&IsNonOverlappingAndDenseIndicator.evalN)r   r   r   r   r   r   r   r   r   r   r     s   r      c                 C   sJ   t | drBzt| W S  ty>   td|  d |  Y S 0 n| S d S )NexpandzRecursionError in sympy.expand())r4   r   r   RecursionErrorr   r   )r   r   r   r   safe_expand	  s    
r   c                 C   s   | | S r*   r   rN   br   r   r   <lambda>  ri   r   c                 C   s   | | S r*   r   r   r   r   r   r     ri   c                 C   s   | | S r*   r   r   r   r   r   r     ri   c                 C   s   | | S r*   r   r   r   r   r   r     ri   c                 C   s
   t | |S r*   )r   r   r   r   r   r     ri   c                 C   s   | |@ S r*   r   r   r   r   r   r     ri   c                 C   s   | |B S r*   r   r   r   r   r   r     ri   c                 C   s
   t | |S r*   )r   r   r   r   r   r     ri   c                 C   s
   t | |S r*   )r"   r   r   r   r   r     ri   )	r   r   r   r   r   andorr   r   c                   C   s   t dd S )Nzshouldn't be hit)rR   r   r   r   r   error"  s    r   c                 C   s   t | tjr^| j}t|dkr^t |d tjr^|d jr^t|d }|d |kr^||d  S t | tjrx| t| kst | tjrt| S || S )Nr   r   r=   )rG   r   ZMulr/   r   r   r   r   )rN   fnZaaZcoefr   r   r   floor_ceil_helper%  s    &&
r   c                 C   s   t | tjS r*   )r   r   r   rM   r   r   r   
floor_impl0  s    r   c                 C   s   t | tjS r*   )r   r   ceilingrM   r   r   r   	ceil_impl3  s    r   c                 C   s   |  S r*   r   rM   r   r   r   r   9  ri   c                 C   s   t | |S r*   )r   Eqr   r   r   r   r   :  ri   c                 C   s   t | |S r*   )r   Ner   r   r   r   r   ;  ri   c                 C   s   t | |S r*   )r   Gtr   r   r   r   r   <  ri   c                 C   s   t | |S r*   )r   Ltr   r   r   r   r   =  ri   c                 C   s   t | |S r*   )r   Ler   r   r   r   r   >  ri   c                 C   s   t | |S r*   )r   Ger   r   r   r   r   ?  ri   c                 C   s   | S r*   r   rM   r   r   r   r   A  ri   c                 C   s   |  S r*   r   rM   r   r   r   r   C  ri   c                 C   s   t | |S r*   )r   Minr   r   r   r   r   D  ri   c                 C   s   t | |S r*   )r   Maxr   r   r   r   r   E  ri   c                 C   s
   t | S r*   )r   rW   rM   r   r   r   r   F  ri   )r   r   r   r   r   r   r   r   r   r   r   r   r   rX   is_non_overlapping_and_densec                  G   s   t |  S r*   )r   r   r   r   r   r   J  ri   )r   r   c                 C   sv   t | }|dkr(|d dkp&| d dk S tt| |tdd}d}|D ](\}}|dkrZqH||krh dS ||9 }qHdS )Nr=   r   r   )keyFT)r   sortedzipoperator
itemgetter)r   rC   rD   Zlengths_and_stridesZexpected_stridelengthstrider   r   r   r   S  s    
r   c                    sb   d  t | |D ]}t|tr|  q(q d us4J t j fdd| D  fdd|D S )Nc                    s   g | ]}t  j|qS r   r`   rI   r   r   r   r   rh   z  ri   z0is_non_overlapping_and_dense.<locals>.<listcomp>c                    s   g | ]}t  j|qS r   r  r   r  r   r   rh   {  ri   )	itertoolschainrG   r   r&   rI   r  )r   rC   r   r   r  r   r  q  s    
>   r   r   r   rX   r   r   >   r   r   r   r   r   >   r   r   r   rX   r   r   r   c                 C   sT   | t v r|  d}n| }| tv r2ttjjj|}n| tv rFtt|}n
tt	|}|S )N_)
2magic_methods_on_operator_with_trailing_underscoremagic_methods_on_submodulegetattrrH   ZfxZexperimentalZsymbolic_shapesmagic_methods_on_mathrV   r  )methodmethod_attrrc   r   r   r   r'     s    
r'   )r   r   r   r   r   r  r  r  Modr"   r   r   r   >   rX   r   r   r   >
   r   r   r   r   r  r   r   r   r   r   c                 C   s\   t | tr| jd ur| jS |  r*t| S |  r:t| S |  rJt| S t	d|  d S )Nzunrecognized return type )
rG   rw   r   r   r   r   r   r   r   rR   r   r   r   r   r&     s    r&   c                    sr   t d  tv r  d}n} fdd} fdd}tv r\ttd| | nttd| | d S )Nr   r  c           	         sJ  t }d }| jd ur.|jd ur.|| j|j}t}|r\|d ur\t| |t| t|S tr~t| t|t| t|fi S t|t	sJ |j
}| j| j
}| j|}z ||}W n2 ty   td d| d| d  Y n0 t|}tv rt}n4tv rt}n$| jtu s.|jtu r4t}n| j}t	|| j||S )Nfailed to eval (, r   )r'   ry    alternate_impl_if_hinted_methodsgetr`   r&   r3   rF   rG   rw   r|   rr   r   r   r   r   r   always_float_magic_methodsrQ   always_bool_magic_methodsrO   r{   )	r,   r   rc   out_hintZalternate_implZ
other_exprr|   outr{   r-   r  r   r   binary_magic_impl  s6    
	

z+_make_node_magic.<locals>.binary_magic_implc              
      s   t }tr$t| t|t| fi S | j| j}z |}W n, tyj   t	
d d| d  Y n0 d }| jd ur|| j}t|}tv rt}ntv rt}n| j}t|| j||S )Nr  r  r   )r'   r3   r`   rF   r&   rr   r   r|   r   r   r   ry   r   always_int_magic_methodsrL   r  rQ   r{   rw   )r,   rc   r|   r   r  r{   r!  r   r   unary_magic_impl  s&    

z*_make_node_magic.<locals>.unary_magic_impl)r
   r  unary_magic_methodssetattrrw   )r  r-   r  r"  r$  r   r!  r   _make_node_magic  s    +r'  c                    s$    fdd}t td | d S )Nc                    s  t tjt }trTt|dd |D dd |D fi }t|tsNJ t||j	S dd |D }dd |D }z g ||R  }W n2 t
y   td d| d| d  Y n0 g }d }	t||D ]}
|
jd u r q||
j q|| }	tt|d	| jt|	S )
Nc                 S   s   g | ]}t |qS r   )r&   r   r   r   r   rh     ri   zH_make_node_sizes_strides.<locals>.sizes_strides_impl.<locals>.<listcomp>c                 S   s   g | ]
}|j qS r   r   r   r   r   r   rh     ri   c                 S   s   g | ]
}|j qS r   r   r   r   r   r   rh     ri   r  z(*z, *r   r=   )r  sysmodulesr   r3   rF   rG   r   rK   rI   r   r   r   r  r  ry   r@   rw   r   r   rr   rO   )r,   r   rC   rc   r   Z
size_exprsZstride_exprsr   hintsr  r   r!  r   r   sizes_strides_impl  s(    $
z4_make_node_sizes_strides.<locals>.sizes_strides_implr  )r&  rw   )r  r-   r+  r   r!  r   _make_node_sizes_strides  s    r,  c                    s   | t v r|  d n|   fdd} fdd} fdd}| tv rZt|d|  d| n0t|d|  d| | tv rt|d	|  d| d S )
Nr  c                    s   t t| j  S r*   )r&   r  rI   r   r  r   r   r$  :  s    z*_make_user_magic.<locals>.unary_magic_implc                    s,   t | j|}|tu rtS tt| j |S r*   r`   rI   r]   r&   r  r,   r   Z
other_noder-  r   r   r"  =  s    z+_make_user_magic.<locals>.binary_magic_implc                    s,   t | j|}|tu rtS tt| | jS r*   r.  r/  r-  r   r   rbinary_magic_implC  s    z,_make_user_magic.<locals>.rbinary_magic_impl__Z__r)r  r%  r&  reflectable_magic_methods)r  Z	user_typer$  r"  r0  r   r-  r   _make_user_magic1  s    r3  c                    s4   t ||  dt|  fdd} j|_|S )z
    Wrapper around lru_cache that clears when new info about shapes has been
    updated.

    Use lru_cache if the output is always the same, regardless of the
    constraints we know now (i.e. evaluate_expr)

    Use _lru_cache otherwise.
    Nc                    s2   |   kr|       | g|R i |S r*   )_get_keycache_clear)r,   r/   r0   Zfn_cacheZ	prior_keyr   r   wrapperg  s    z_lru_cache.<locals>.wrapper)r
   	functoolswraps
cache_info)r   maxsizer7  r   r6  r   
_lru_cacheZ  s    
r<  c                       sJ   e Zd ZU ddgZee ed< ee ed< ee ed<  fddZ	  Z
S )Symbolsourcesstack	__slots__c                    s*   t  j| g|R i |}g |_d |_|S r*   )super__new__r>  r?  )r   r/   r0   r,   	__class__r   r   rB  }  s    zSymbol.__new__)r   r   r   r@  r   r   __annotations__r   r   rB  __classcell__r   r   rC  r   r=  x  s   
r=  c                       s*   e Zd Z fddZedddZ  ZS )ShapeGuardPrinterc                    s   t    || _|| _d S r*   )rA  r   symbol_to_source
source_ref)r,   rH  rI  rC  r   r   r     s    
zShapeGuardPrinter.__init__r   c                 C   s^   t |tsJ tt||| jv sJJ | ddd |jD  d| j | | j| d S )Nz (could be from c                 S   s   g | ]}|  qS r   namer   r   r   r   rh     ri   z3ShapeGuardPrinter._print_Symbol.<locals>.<listcomp>z	) not in r   )rG   r=  r   rK   rH  r>  rI  )r,   r|   r   r   r   _print_Symbol  s    zShapeGuardPrinter._print_Symbol)r   r   r   r   r   rL  rF  r   r   rC  r   rG    s   	rG  c                   @   sz  e Zd Zdd Zdd Zedd Zdd Zej	e
d	d
dZdee dddZdd Zdd Zee
ddddZdd Zdd fddee dddZd d! Zd"d# Zd$d% ZdEd&d'Zd(d) Zedd*d+d,d-Zeddd+d.d/Zed0d1 Zeddd+d2d3Ze d4dd5d6d7Z!d8d9 Z"ed:dd;d<d=Z#e d4e$d> e%d?d@dAdBZ&e d4dFdd5dCdDZ'd?S )Gr    c                 C   sJ   g | _ i | _i | _t | _tdtdd| _t	 | _
t	 | _d S )Nr   r=   r   r=   )guardsr~   r   set	divisibler   r   
val_to_varr  countunbacked_symfloat_counterunbacked_symint_counterr   r   r   r   r     s    
zShapeEnv.__init__c                 C   s   t tddS )Nsuppress_guardsF)r  TLSr   r   r   r   _suppress_guards_tls  s    zShapeEnv._suppress_guards_tlsc                 c   s$   dt _zd V  W dt _ndt _0 d S )NTF)rV  rU  r   r   r   r   rU    s    zShapeEnv.suppress_guardsc                 C   s   t | jt | jfS )z
        Defines the current "state" of the guards we've accumulated in this ShapeEnv.
        Determines when we need to invalidate our cache
        )r   r   rP  r   r   r   r   r4    s    zShapeEnv._get_key)exsourcec                    s  ddl mm   fddt D dgt t D ]\}}|dv rJt||< qJt	dd D rrfd	d
t
tD }tfddt
tD }|D ]^\}}| du r | |v r| |  |< | |  || |  < qt	dd D rjtfddt
tD \}}| j||< qjtdd D sJ fddt D }g }	tD ]2\}}
|
dusJ |	j|
|d qj  j d}||	|fS )z
        Returns a list of symbolic sizes and strides for the given tensor.
        We try our best to express stride in terms of the sizes, so as to not
        introduce new symbolic variables.
        r   )TensorPropertySourceTensorPropertyc              	      s&   g | ]\}} | j|qS r   )create_symbolSIZE)rf   irb   )r[  rZ  r,   rY  r   r   rh     s   zIShapeEnv.create_symbolic_sizes_strides_storage_offset.<locals>.<listcomp>NrM  c                 s   s   | ]}|d u V  qd S r*   r   rf   r   r   r   r   r     ri   zHShapeEnv.create_symbolic_sizes_strides_storage_offset.<locals>.<genexpr>c                    sL   i | ]D}| d ur   | dkr |   |  | |  qS r   )r  sizerf   r^  )rX  r`  r  r   r   
<dictcomp>  s   zIShapeEnv.create_symbolic_sizes_strides_storage_offset.<locals>.<dictcomp>c                    s(   g | ] }| d u r   | |fqS r*   r  ra  rX  r  r   r   rh     ri   c                 s   s   | ]}|d u V  qd S r*   r   r_  r   r   r   r     ri   c                    s(   g | ] }| d u r   | |fqS r*   rc  ra  rd  r   r   rh     s   c                 s   s   | ]}|d uV  qd S r*   r   r_  r   r   r   r     ri   c                    s   g | ]\}} j ||d qS )rx   )create_symintnode)rf   r^  ry   r   r   r   rh     ri   rx   )torch._dynamo.sourcerZ  r[  	enumerater`  r   r  r   r   anyranger  minr\  STRIDEr   r  r@   re  storage_offsetSTORAGE_OFFSET)r,   rX  rY  r^  rb   
candidatesval_listr  Zsym_sizeZ
sym_strideZstride_exprZsym_storage_offsetr   )r[  rZ  rX  r,   r`  rY  r  r   ,create_symbolic_sizes_strides_storage_offset  sT    

(


z5ShapeEnv.create_symbolic_sizes_strides_storage_offsetz
sympy.Expr)symry   c                C   s   t t|| t|S r*   )r   rw   rL   )r,   rq  ry   r   r   r   re    s    zShapeEnv.create_symintnodec                 C   sD   t dt| j }dtt d d |_tt	|| t
d S )NfrT   r>   )r=  nextrS  join	tracebackformat_listextract_stackr?  r   rw   rQ   r,   symbolr   r   r   create_unbacked_symfloat  s    z!ShapeEnv.create_unbacked_symfloatc                 C   sH   t dt| j dd}dtt d d |_tt	|| t
d S )Nr^  T)integerrT   r>   )r=  rs  rT  rt  ru  rv  rw  r?  r   rw   rL   rx  r   r   r   create_unbacked_symint  s    zShapeEnv.create_unbacked_symint)rb   rY  r   c                 C   s   t |ts J t| d| ts,td|dk rTddlm} | | || S || jvrt	dt
| j ddd}t|| j|< || j|< | |}t |t	r|j| |S )N z.Need sympy installed to create symbolic shapesr   )NegateSourcer   TZpositiver{  )rG   r   rK   	HAS_SYMPYr5   rf  r~  r\  rQ  r=  r   r~   r   r   duck_intr>  r@   )r,   rb   rY  r~  Z
sympy_exprr   r   r   r   r\  
  s     



zShapeEnv.create_symbolc                 C   s2   || j v s(J d| dt| j   | j | S )NzDirect call to duck_int MUST only duck size an integer values that have already produced by inputs (allocated by create_symbol), or we risk being unable to instantiate the symbolic variable later.  However, at time of this call val=z1 was not duck sized.  Bound duck sized integers: )rQ  rA   r   )r,   rb   r   r   r   r  *  s    zShapeEnv.duck_intc                 C   s   |   S r*   rJ  )rg   r   r   r   r   E  ri   zShapeEnv.<lambda>F)_simplifiedr   c             	      s6  ddl m m}m} g tt fdd}t||D ]\}}	t|	t	rbddl m
}
 |
|	}	t|	tspJ |d u rzq<t|tr||	| q<t|tjsJ t| D ]\}}|||	|j|| qt| D ]\}}|||	|j|| q|||	|j|  q<g }|szD ]^\}	}t|trN|v rN|	| d krNqt||}|||	 d|  q| jD ]j\}}| |d urq| |}z|t|| W n& ty   td|   Y n0 q|s2  D ]6}|sJ |||d  d||d  d	 q|S )
Nr   )r~  rZ  r[  c                    sx   t |tr`|jj}t |tjr.| |  n"t | tjrP|   |  | |f n| t|f d S r*   )rG   r   rI   r|   r   r=  r@   r   )rY  rb   r   r~  Zinput_guardsrH  r   r   track_symint  s    
z-ShapeEnv.produce_guards.<locals>.track_symint)LocalSourcez == zFailing guard allocated at: 
z
 != 0 and z != 1)!rf  r~  rZ  r[  collectionsdefaultdictrA   r  rG   r   r  r   r   rH   Tensorrg  r`  r]  r  rk  rm  rl  r=  rG  Zdoprintr@   rN  _maybe_evaluate_staticr   r   r   r   values)r,   placeholdersr>  rI  r  rZ  r[  r  trY  r  r^  r   exprsr|   Zsexprgtbr   r  r   produce_guardsD  s\    =





*zShapeEnv.produce_guardsc                    sd   ddl m  dd tt|D }| | fdd|D }|r`d|}t|i tt||S dS )Nr   GlobalSourcec                 S   s   g | ]}d | qS )r  r   ra  r   r   r   rh     ri   z5ShapeEnv.evaluate_guards_for_args.<locals>.<listcomp>c                    s   g | ]} |qS r   r   r   r  r   r   rh     ri   z and T)	rf  r  ri  r   r  rt  r   dictr  )r,   r  r/   	arg_namesrN  coder   r  r   rs     s    
z!ShapeEnv.evaluate_guards_for_argsc                    s   i   fdd}t ||D ]\}}|d u r,qt|trB||| qt|tjsRJ t| D ]\}}|||| q^t| D ]\}}|||| q|| |  q S )Nc                    s   t |tr|jj}t |tjrT| v rJ | | ksRJ  |  d|  q|  |< nPt | tjr|  v r |  |  ksJ  |   d|   n|   | < d S )Nz != )rG   r   rI   r|   r   r=  )argrb   r   Zbindingsr   r   bind_symint  s    
$

,z*ShapeEnv.bind_symbols.<locals>.bind_symint)	r  rG   r   rH   r  rg  r`  r  rl  )r,   r  r/   r  r  r  r^  r   r   r  r   rv     s    

zShapeEnv.bind_symbolsc                    s    fdd j D S )Nc                    s(   g | ] }  |jd u r |jqS r*   )r  r|   r   rf   Zguardr   r   r   rh     ri   z2ShapeEnv.get_nontrivial_guards.<locals>.<listcomp>)rN  r   r   r   r   get_nontrivial_guards  s    zShapeEnv.get_nontrivial_guardsc                    s&   fdd d  fdd| jD S )Nc                    s    sdS dt | d S )NrT   z
   Guarded at:
z   )textwrapindent)r  )verboser   r   	format_tb  s    z)ShapeEnv.format_guards.<locals>.format_tb
c                 3   s$   | ]}d |j   |j V  qdS ) - N)r|   r?  r  )r  r   r   r     ri   z)ShapeEnv.format_guards.<locals>.<genexpr>)rt  rN  )r,   r  r   )r  r  r   format_guards  s    zShapeEnv.format_guardsc                 C   s0   t t}| j D ]\}}|| | q|S r*   )r  r  rA   r   itemsr@   )r,   Zshape_groupskvr   r   r   get_shape_groups  s    
zShapeEnv.get_shape_groupszOptional[sympy.Expr])r|   r   c                    s     |}t|j} fddt|D }||}i }|tD ]"}t|j	d |j	d  ||< qBt
||}tt|jdkr|S dS )zC
        Tries to evaluate expr without introducing guards
        c                    s6   i | ].\}}| j v r|tjd | dddd qS )Zshape_Tr  r=   )r~   r   r=  )rf   idxr  r   r   r   rb  +  s   
z3ShapeEnv._maybe_evaluate_static.<locals>.<dictcomp>r   r=   N)r   rA   r   rg  r}   atomsr"   r   r   r/   r   r   )r,   r|   symbolsZnew_shape_envZnew_exprZfloor_div_replaceatomr   r   r   r  #  s    



 zShapeEnv._maybe_evaluate_staticc                    s"    fdd|j D }t||S )Nc                    s    i | ]}|  ttj|qS r   )_findr   r   r=  r   r   r   r   rb  <  ri   z$ShapeEnv.replace.<locals>.<dictcomp>)r   r   r}   )r,   r|   r   r   r   r   r   :  s    zShapeEnv.replacec                 C   s>   t  }| jD ]&}| |}t|jdkr|| q|| _d S r   )rO  rP  r   r   r   r   )r,   Znew_divisibler  resr   r   r   _update_divisible?  s    

zShapeEnv._update_divisiblec                 C   sv   |  |}|trr|   i }|tD ]4}|j\}}|  || | jv r*t|| ||< q*|	|}t
|}|S r*   )r   hasr"   r  r  r/   rP  r   r   r}   r   )r,   r|   Zdiv_replacementsr  r   r   r   r   r   r   I  s    



zShapeEnv.simplifyr   r   c                 C   s,   t || j}t|jdkr(| ||S )a  
        Gets a size hint for a given expression from the underlying shapes we had.
        Does not introduce a guard, so only use this when you can guarantee that
        your code is still valid for arbitrary shapes (such as optimization decisions)
        r   )r   r}   r~   r   r   r   )r,   r|   Zresult_exprr   r   r   	size_hintW  s    
zShapeEnv.size_hintc                 C   s,   d dd |jD }td| d| dS )Nz

c                 s   s    | ]}d | d|j  V  qdS )zData dependent variable 'z' allocated at:
N)r?  r   r   r   r   r   f  s   z6ShapeEnv._make_data_dependent_error.<locals>.<genexpr>z
GuardOnDataDependentSymNode: It appears that you're trying to get a value out of symbolic int/float whose value is data-dependent (and thus we do not know the true value.)  The expression we were trying to evaluate is zT.  Scroll up to see where each of these data-dependent accesses originally occurred.)rt  r   r   )r,   r|   Zaccessesr   r   r   r   c  s    
z#ShapeEnv._make_data_dependent_errorzsympy.Symbol)rN   r   c                    sL   | j vr|S  j | } fdd|jD } j | | j |<  j | S )z
        Implements a DSU-like algorithm to find the variable that represents a
        Also handles transitive non-identity replacements.

        a: b + c
        c: d
        c                    s   i | ]}|  |qS r   )r  r   r   r   r   rb    ri   z"ShapeEnv._find.<locals>.<dictcomp>)r   r   r}   )r,   rN   r  Zcur_replacer   r   r   r  u  s    	

zShapeEnv._find)zsympy.Eqzsympy.NeN)r|   concrete_boolr   c           
         s  t |tu sJ t|tjr&|s:dS nt|tjr:|r:dS t|j}t|dksXJ dt|dkrhdS t	| fdddd}|j
}|j}|tjsZzxtj|| |d dd	}t|d
krW dS |d |d  }tdd t|D r |}| jttj|d < W nH ty$   Y n6 tyX   td| d| d|d  d Y n0 |tjrt|tjd }	zDtj|| |	dd	}t|d
kr|d |	 dkr j|	 W n ty   Y n0 dS )z
        Evaluates the result of an eq call. If true, uses information to
        simplify shapes (i.e. a == b or a % 5 == 0)
        Nr   z1The expression should not be static by this point   c                    s     | | jfS r*   )r  rK  r   r   r   r   r     ri   z*ShapeEnv._maybe_guard_eq.<locals>.<lambda>T)r  reverse)r  r=   c                 s   s   | ]}|j V  qd S r*   )r   )rf   r  r   r   r   r     ri   z+ShapeEnv._maybe_guard_eq.<locals>.<genexpr>zRecursionError in sympy.solve(r  r  r   )rK   rO   rG   r   r   r   rA   r   r   r  lhsrhsr  r  Zsolver   Zpreorder_traversalr  r   r   r=  r+   r   r   r   tupler  rP  r   )
r,   r|   r  freer  r  Z	solutionsZsolutionZnew_varZmod_exprr   r   r   _maybe_guard_eq  sH    

( zShapeEnv._maybe_guard_eqc                 C   s   t |jdkr|S | |}| |}|dur2|S |du rF| |}n
t|}t|tjtj	frr| 
|t| |  sdtt dd }|tju r| jt|| n>|tju r| jtt|| n| jtt||| |S )zO
        Given an expression, evaluates it, adding guards if necessary
        r   NrT   )r   r   r   r  r  r   ZsympifyrG   r   r   r  rO   rW  rt  ru  rv  rw  r   rN  r@   r   r   Not)r,   r|   ry   Zstatic_exprZconcrete_valr?  r   r   r   r     s*    


	

zShapeEnv.evaluate_expr)F)N)(r   r   r   r   rW  r	   rU  r4  rH   r  r   rp  r   rL   re  rz  r|  r\  r  r   r   r  rs   rv   r  r  r  r<  r  r   r  r   r
   r  r   r  r   rO   r  r   r   r   r   r   r      sP   
=  	/

	0r    )N){rH   typingr   r   r   r   r   r   r   r(  builtinsr  r  rV   r8  	threading
contextlibr	   r
   ru  r  r  loggingr   r   r   r   r   r   r   Ztorch._guardsr   r   rY   	getLoggerr   r   r5   r   r   Zsympy.printing.precedencer   Zsympy.printing.strr   Zsympy.core.logicr   r   r  ImportErrorZ_opsopsZaten__all__r3   r!   r   r   rF   r(   r%   rP   r#   r$   rX   r`   ro   rq   ru   rv   rw   ZFunctionr   r   r"   r   r   r2  r   r   r   r   Zmagic_methodsZsizes_strides_methodsrj  maxr  r   r  r%  Zbool_magic_methodsr  r  r  r'   r   r   r   r   r   r   r   r   r   r   r   r)   r  r#  r  r&   r'  r,  r  r  r-   r3  r<  ZDummyr=  rG  localrV  r    r   r   r   r   <module>   s  $$




 aM
	S


