a
    d3                     @   sr  d dl Z d dlmZmZmZmZmZmZ d dlmZ	 d dl
mZmZ d dlmZmZmZmZmZ d dlZd dlmZ d dlZd dlmZ eee jeddd	Zeee jedd
d	Zeee jeddd	Zede jdddd	Zdd	 ZeeedddZdd ZG dd dZeedddZddeee dddZ!dde"e dddZ#d d! Z$eed"d#d$Z%dS )%    N)Number
NumberType
TensorLikeTensorLikeType	ShapeTypeELEMENTWISE_TYPE_PROMOTION_KIND)tree_flattentree_unflatten)CallableSequenceTuple
NamedTupleoverload)wraps)chain)adtypereturnc                 C   s   d S N r   r   r   r   e/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_prims_common/wrappers.py_maybe_convert_to_dtype   s    r   c                 C   s   d S r   r   r   r   r   r   r      s    c                 C   s   d S r   r   r   r   r   r   r      s    c                 C   s   d S r   r   r   r   r   r   r      s    c                    s|   t | tr"| j kr|  S | S t | tr:t | S t | trZt fdd| D S | d u rfd S t	d
t| d S )Nc                 3   s   | ]}t | V  qd S r   r   .0xr   r   r   	<genexpr>,       z*_maybe_convert_to_dtype.<locals>.<genexpr>z7Received type {0} that is neither a tensor or a number!)
isinstancer   r   tor   utilsZdtype_to_type_ctorr   tuple
ValueErrorformattyper   r   r   r   r   $   s    




)r   typr   c                 C   sR   t | ts dt| }t|tt| |sJd| t| |}t||| S )Nz6Found unknown type {0} when trying to convert scalars!z9Scalar {0} of type {1} cannot be safely cast to type {2}!)r    r   r%   r&   r$   r"   Zis_weakly_lesser_type)r   r'   msgr   r   r   _maybe_convert_to_type7   s    

r)   c                 C   s4   t |dr,|jD ]}t| |dr dS qdS | |u S )N__args__)r'   
annotationTF)hasattrr*   _annotation_has_type)r'   r+   r   r   r   r   r-   D   s    

r-   c                   @   s:   e Zd ZdZddeee dddZeeddd	Z	dS )
"elementwise_type_promotion_wrappera  
    Adds elementwise type promotion to a Python reference implementation.

    Takes two kwargs, type_promoting_args and type_promotion_kind.

    type_promoting_args must be a string Sequence specifiying the argument names of all
    arguments that participate in type promotion (and should be type promoted). If the
    arg specifies a Sequence-type then every element of the Sequence will participate in
    type promotion.

    type_promotion_kind must be one of the kinds specified by ELEMENTWISE_TYPE_PROMOTION_KIND.
    See its documentation for details.

    Other type promotion behavior, like validating the Python type of scalar arguments, must
    be handled separately.
    N)type_promoting_args)type_promotion_kindr/   c                C   s   || _ || _d S r   )type_promoting_arg_namesr0   )selfr0   r/   r   r   r   __init__`   s    z+elementwise_type_promotion_wrapper.__init__fnr   c                    s,   t  t  fdd}|_|S )Nc                     s   j | i | t fddjD }t|d }tj|dji\ fddjD } j| f i  j}t	|t
rt|S t	|trtfdd|D S tdt| d S )	Nc                 3   s&   | ]}| j  v r j | V  qd S r   )	argumentskeysr   )boundr   r   r   o   s   zKelementwise_type_promotion_wrapper.__call__.<locals>._fn.<locals>.<genexpr>r   r0   c                    s,   i | ]$}| j  v r|t j | qS r   )r6   r7   r   r   )r8   compute_dtyper   r   
<dictcomp>{   s   zLelementwise_type_promotion_wrapper.__call__.<locals>._fn.<locals>.<dictcomp>c                 3   s   | ]}t | V  qd S r   r   r   )result_dtyper   r   r      r   zUnhandled result type: )bindr#   r1   r   r"   Zelementwise_dtypesr0   r6   updater    r   r   r   AssertionErrorr&   )argskwargsr/   Zflattened_type_promoting_argsZpromoted_argsresultr5   r2   sig)r8   r9   r;   r   _fnl   s(    



z8elementwise_type_promotion_wrapper.__call__.<locals>._fninspect	signaturer   __signature__)r2   r5   rD   r   rB   r   __call__i   s
    
z+elementwise_type_promotion_wrapper.__call__)
__name__
__module____qualname____doc__r   r   strr3   r
   rI   r   r   r   r   r.   N   s   	r.   )outshapec                 C   sH   t | j|r| S |  dkr:dt| j d}t| | |S d S )Nr   zCAn output with one or more elements was resized since it had shape a   which does not match the required output shape {str(shape)}. This behavior is deprecated, and in a future PyTorch release outputs will not be resized unless they have zero elements. You can explicitly reuse an out tensor t by resizing it, inplace, to zero elements with t.resize_(0).)r"   Z
same_shaperP   ZnumelrN   warningswarnZresize_)rO   rP   r(   r   r   r   _maybe_resize_out   s    
rS   F)exact_dtype	copy_fromcopy_torT   c                    sv    j j kr$d j j }t||rHt jjk fdd n$ttj jjd fdd  S )NzZAttempting to copy from device {0} to device {1}, but cross-device copies are not allowed!c                      s   d j  dj  dS )Nz"Expected out tensor to have dtype z	 but got z insteadr   r   rV   rW   r   r   <lambda>   s   
z _safe_copy_out.<locals>.<lambda>)Z	cast_fromZcast_toc                      s   d j  dj  dS )NzAttempting to cast from z to out tensor with dtype z0, but this can't be cast because it is not safe!r   r   rX   r   r   rY      r   )devicer%   RuntimeErrorr"   checkr   Zcan_safe_cast_toZcopy_)rV   rW   rT   r(   r   rX   r   _safe_copy_out   s    
r]   )	out_namesrT   c                    s<   t dks t dks J ttd fdd}|S )Nr      r4   c              	      s   rt nttdd ttD  }r.t ntdj dd D td t	fdd D t
dd	 fd
d
}tjdtjjd|d}jj|fv sJ tj |f}tj|d|_j|_||jd< |jd< |S )z?
        Adds the out parameter to a Python reference.
        c                 s   s   | ]
}t V  qd S r   r   )r   _r   r   r   r      r   z4out_wrapper.<locals>._out_wrapper.<locals>.<genexpr>Zreturn_types_c                 S   s   g | ]}|t fqS r   r`   )r   or   r   r   
<listcomp>   r   z5out_wrapper.<locals>._out_wrapper.<locals>.<listcomp>)rZ   r   c                 3   s   | ]}| j v V  qd S r   )
parameters)r   p)rC   r   r   r      r   N)rO   c                    s  r0 d ur0D ]}t  |}||vr|||< q|i |ttrLsjttrfttksjJ  d urrt tsJ t j t d n`t tsJ t	t tk fddt
 t D ]"\}}t||j t||d qn r S   S )NrU   c                      s   dt  dt   S )Nzexpected tuple of z elements but got )lenr   rO   rA   r   r   rY      r   z@out_wrapper.<locals>._out_wrapper.<locals>._fn.<locals>.<lambda>)getattrr    r   r   rf   rS   rP   r]   r"   r\   	TypeErrorzip)rO   r?   r@   kZout_attrrrb   )rT   factory_kwargsr5   is_factory_fn	is_tensorr^   return_typerg   r   rD      s<    


z.out_wrapper.<locals>._out_wrapper.<locals>._fnrO   )kinddefaultr+   )rd   return_annotationr   )r   r   r#   rangerf   r   rJ   rF   rG   allr   	ParameterKEYWORD_ONLYrs   emptyr   rd   values	SignaturerH   __annotations__)r5   Zout_typerD   Z	out_paramparamsrT   ro   r^   )rm   r5   rn   rp   rC   r   _out_wrapper   s<    
 4

z!out_wrapper.<locals>._out_wrapper)rf   r
   )rT   r^   r~   r   r}   r   out_wrapper   s    [r   c                    s>   fddG fdddt jj t fdd}|S )Nc                    s8   t j }z$t jt jjj} | i |W ~S ~0 d S r   )torchZ_CZ_AutoDispatchBelowAutogradZ&_dispatch_tls_is_dispatch_key_excludedZDispatchKeyZADInplaceOrView)r?   r@   gold)primr   r   redispatch_prim  s    
z0backwards_not_supported.<locals>.redispatch_primc                       s(   e Zd Ze fddZedd ZdS )z6backwards_not_supported.<locals>.BackwardsNotSupportedc                    s   t ||\}} ||S r   )r	   )ctx	args_spec	flat_argsr?   r@   r   r   r   forward'  s    z>backwards_not_supported.<locals>.BackwardsNotSupported.forwardc                 W   s   t dd S )Nzbackwards not supported on prim)r[   )r   r?   r   r   r   backward,  s    z?backwards_not_supported.<locals>.BackwardsNotSupported.backwardN)rJ   rK   rL   staticmethodr   r   r   r   r   r   BackwardsNotSupported&  s   r   c                     sJ   t | |f\}}t r<tdd |D r< j|g|R  S | |S d S )Nc                 s   s    | ]}t |tjr|jV  qd S r   )r    r   TensorZrequires_grad)r   r   r   r   r   r   3  r   zBbackwards_not_supported.<locals>._autograd_impl.<locals>.<genexpr>)r   r   Zis_grad_enabledanyapply)r?   r@   r   r   )r   r   r   r   _autograd_impl0  s    z/backwards_not_supported.<locals>._autograd_impl)r   ZautogradZFunctionr   )r   r   r   )r   r   r   r   backwards_not_supported  s
    
r   r4   c                    s(   t  }t  fdd}||_|S )zQ
    Allows unary operators that accept tensors to work with Python numbers.
    c                     s~   t | dkrpt| d trptt| d }t| }tj| d |d|d<  |i |}t|tj	shJ |
 S  | i |S )Nr   r   )rf   r    r   r"   Ztype_to_dtyper&   listr   Ztensorr   item)r?   r@   r   Zargs_rA   r5   r   r   rD   N  s    z-elementwise_unary_scalar_wrapper.<locals>._fnrE   )r5   rC   rD   r   r   r    elementwise_unary_scalar_wrapperH  s
    
r   )&r   Ztorch._prims_commonr   r   r   r   r   r   Z_prims_commonr"   Ztorch.utils._pytreer   r	   typingr
   r   r   r   r   rF   	functoolsr   rQ   	itertoolsr   r   r   r&   r)   r-   r.   rS   boolr]   rN   r   r   r   r   r   r   r   <module>   s8    
Ab+