a
    d[                  C   @   s  U d dl mZmZmZmZ d dlZd dlmZ d dlm	Z	m
Z
mZmZmZmZmZmZ d dlmZmZmZ dZejedZejeddZejedd	Zejedd
ZejeddZg dZi Zeeef e d< h dZ!edddZ"e!D ],Z#e$de# de# de# de# de# d qh dZ%e%D ].Z#e$de# de# de# de# de# d q,dhZ&e&D ].Z#e$de# de# de# de# de# d qfdd Z'eeeed d!d"Z(eeej)d#d$d%Z*d&d' Z+d(d) Z,d*d+ Z-d,d- Z.eee	d.d/d0Z/eee	e0d1d2d3Z1dgeee	ee2 e2e0d5d6d7Z3eed8d9d:Z4eee	d.d;d<Z5eee	d.d=d>Z6dd?eed@dAdBZ7dddd4d4dCeeeeej) eej8 eej9 e2e2dDdEdFZ:e'edG< e(edH< e*edI< e7edJ< e+edK< e,edL< e-edM< e.edN< e4edO< e/edP< e1edQ< e3edR< e5edS< e6edT< e:edU< dVdW Z;d4d4d4d4d4d4d4dXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdXdYBZ<eee2f e dZ< d[d\ Z=d]d^ Z>d_d` Z?dadb Z@dcdd ZAdedf ZBdS )h    )AnyDictOptionalTupleN)DimsSequenceTypeelementwise_dtypesELEMENTWISE_TYPE_PROMOTION_KINDgetnvFuserDtypemake_contiguous_strides_for
NumberType	ShapeTypeTensorLikeType)_maybe_convert_to_dtypebackwards_not_supported"elementwise_type_promotion_wrappernvprimsZDEFZIMPLZCompositeExplicitAutogradZCompositeImplicitAutogradZAutogradZMeta)=absacosasinatanatanhcoscoshclonebitwise_notceilerferfcexpexpm1floorimagisfinitelgammaloglog1plog2log10real
reciprocalnegroundrsqrtsignsinsinhsqrttantanh	transposetruncaddatan2bitwise_and
bitwise_orbitwise_xordiveqfmodgegtleltmulnepow	remaindersubsqueezeview_ofbroadcast_in_dimwhereconvert_element_typesumvaramaxamin_nvfuser_impls>!   r&   r   r0   r   r$   r   r/   r   r   r)   r   r   r%   r'   r*   r   r!   r(   r"   r2   r,   r   r   r   r1   r4   r.   r    r   r+   r#   r   r-   )fnamec                 C   s8   z ddl m} t|j| sJ W n ty2   Y n0 d S )Nr   )FusionDefinition)
nvfuser._CrQ   getattrZ	OperatorsImportError)rP   fd rV   c/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_prims/nvfuser_prims.py_assert_nvfuser_op_exists   s
    rX   zL
# Ensure that the nvfuser implementation exists
_assert_nvfuser_op_exists("z	")

def _z#_nvfuser(fd, a):
    return fd.ops.z3(a)  # type: ignore[attr-defined]

_nvfuser_impls["z"] = _z	_nvfuser
>   r5   r?   rB   r7   r;   rC   rE   r6   rD   rA   r<   r@   r>   r:   r=   r9   r8   z&_nvfuser(fd, a, b):
    return fd.ops.z6(a, b)  # type: ignore[attr-defined]

_nvfuser_impls["rI   z)_nvfuser(fd, a, b, c):
    return fd.ops.z9(a, b, c)  # type: ignore[attr-defined]

_nvfuser_impls["c	           	   
   C   s   | j ||||||||S )a  
    if weight is None:
        weight = fd.define_null_tensor()
    if bias is None:
        bias = fd.define_null_tensor()
    if running_mean is None:
        running_mean = fd.define_null_tensor()
    if running_var is None:
        running_var = fd.define_null_tensor()
    )opsZ
batch_norm)	rU   inputweightbiasrunning_meanrunning_vartrainingmomentumepsrV   rV   rW   _native_batch_norm_nvfuser   s    rb   rU   ashapeZbroadcast_dimensionsc                 C   s   | j |||S N)rY   rH   rc   rV   rV   rW   _broadcast_in_dim_nvfuser   s    rg   )rU   rd   dtypec                 C   s   t |}| j||S rf   )r	   rY   cast)rU   rd   rh   nvfuser_dtyperV   rV   rW   _convert_element_type_nvfuser   s    rk   c                 C   s   | j ||S rf   )rY   ZpermuterU   rd   dimsrV   rV   rW   _transpose_nvfuser  s    rn   c                 C   sB   t |ddD ]0}| j|||}|d | ||d d   }q|S )NT)reverse   )sortedrY   rF   )rU   rd   a_shapeZ
dimensionsidxrV   rV   rW   _squeeze_nvfuser  s    rt   c                 C   s   | j |S rf   rY   setrU   rd   rV   rV   rW   _view_of_nvfuser  s    rx   c                 C   s   | j |||S rf   )rY   view)rU   rd   rr   	new_shaperV   rV   rW   _view_nvfuser  s    r{   rl   c                 C   s(   d}ddl m} |j}| j||||S )NFr   )DataType)rR   r|   ZNullrY   rK   )rU   rd   rm   	keep_dimsr|   output_dtyperV   rV   rW   _sum_nvfuser  s    r   )rU   rd   rm   
correctionc                C   s   d}| j ||||S NF)rY   rL   )rU   rd   rm   r   r}   rV   rV   rW   _var_nvfuser&  s    r   FrU   rd   rm   unbiasedkeepdimr   c                C   s"   |d u sJ d}| j ||||S r   )rY   var_meanr   rV   rV   rW   _var_mean_nvfuser1  s    
r   rw   c                 C   s   | j |S rf   )rY   	rand_likerw   rV   rV   rW   _rand_like_nvfuserB  s    r   c                 C   s   d}| j |||S r   )rY   maxrU   rd   rm   r}   rV   rV   rW   _amax_nvfuserF  s    r   c                 C   s   d}| j |||S r   )rY   minr   rV   rV   rW   _amin_nvfuserO  s    r   )memory_format)rU   rZ   c                C   s   | j |S rf   ru   )rU   rZ   r   rV   rV   rW   _clone_nvfuserX  s    r   )rh   layoutdevice
pin_memoryrequires_grad)rU   re   
fill_valuerh   r   r   r   r   c          	      C   sr   |t dksJ |d u s(|t ju s(J |du s4J |du s@J |d urL|ntt|}t|}| j|||S )NcpuF)	torchr   ZstridedutilsZtype_to_dtypetyper	   rY   full)	rU   re   r   rh   r   r   r   r   rj   rV   rV   rW   _full_nvfuser\  s    r   native_batch_normrH   rJ   r   r3   rF   rG   ry   r   rK   rL   r   rM   rN   r   c                  C   s   d} t d d d d d ddddd}d d d d ddddd}t| | t| | ttjjj	| }|j
}t| t| ||fD ]*}d	|_td |_td |_tjjj|_qd S )
Nr   zfull(SymInt[] size, Scalar fill_value, *, ScalarType? dtype=None, Layout? layout=None, Device? device=None, bool? pin_memory=None, bool? requires_grad=None) -> TensorFoutrh   r   r   r   r   c          	      S   s   t | }tjjd | |||dS N)re   stridesrh   r   )r
   r   _prims
TensorMeta)	sizer   r   rh   r   r   r   r   r   rV   rV   rW   
_meta_impl  s    z!register_full.<locals>._meta_implc             
   S   s   t j| |||||||dS )Nr   )r   r   )r   r   r   rh   r   r   r   r   rV   rV   rW   
_prim_impl  s    z!register_full.<locals>._prim_implz5Create a tensor with given size and filled with valuenvprimdefinenvprim_implimplnvprim_meta_implrS   r   _opsrY   r   defaultnvprim_autograd_implr   __doc__rO   impl_nvfuser_nvfuser_is_recomputableis_recomputable_prims_commonRETURN_TYPENEWreturn_type)namer   r   prim_packetprimprV   rV   rW   register_full  s8    	

r   T)BrM   rN   rK   rL   r   r   r   r   r   r5   r   r   r6   r   r7   r   r8   r9   rH   r   r   rJ   r   r   r:   r;   r   r   r   r   r    r<   r   r=   r>   r!   r"   r?   r#   r$   r'   r%   r&   r@   rA   rB   r*   rC   r(   r)   rD   r+   r,   r-   r.   r/   r0   rF   rE   r1   r2   r3   r4   ry   rG   rI   r   c                     s  d} t |  dd d  dd }t| | tjjjj}|j	tj
ttj
 ttj
 ttj
 ttj
 tttttj
tj
tj
f d	fdd	 tj
ttj
 ttj
 ttj
 ttj
 tttttj
tj
tj
f d	 fd
d}t| | |fD ]*}d|_td |_td |_tjjj|_qdS )z]This function is used to register the native_batch_norm function in torch.ops.nvprims module.r   zX(Tensor input, Tensor? weight, Tensor? bias, Tensor? running_mean, Tensor? running_var, z)bool training, float momentum, float eps)z -> (Tensor, Tensor, Tensor)c              
   S   s   t | |||||||S rf   )r   r   )rZ   r[   r\   r]   r^   r_   r`   ra   rV   rV   rW   r     s    z.register_native_batch_norm.<locals>._prim_impl)	rZ   r[   r\   r]   r^   r_   r`   ra   returnc              	      sl   t j| jrtd| j}t| ||tjd\}	}
t| |	} ||||||||\}}}t||}|||fS )N!Complex tensors are not supported)type_promotion_kind)	r   r   is_complex_dtyperh   NotImplementedErrorr   r   Z	NO_OPMATHr   )rZ   r[   r\   r]   r^   r_   r`   ra   Zresult_dtypeZcomputation_dtype_Zinput_outputmeanZrstdZoutput_r   rV   rW   _native_batch_norm_ref#  s    



z:register_native_batch_norm.<locals>._native_batch_norm_refc           	   
      sP   ddl m} | * t | |||||||W  d    S 1 sB0    Y  d S )Nr   NvfuserPrimsModeZtorch._prims.contextr   r   )	rZ   r[   r\   r]   r^   r_   r`   ra   r   )r   rV   rW   _native_batch_norm_autogradA  s
    z?register_native_batch_norm.<locals>._native_batch_norm_autogradzComputes batch normalization.N)r   r   r   r   r   r   rY   r   r   r   ZTensorr   boolfloatr   r   r   rO   r   r   r   r   r   r   r   )r   r   r   r   r   rV   )r   r   rW   register_native_batch_norm  sL    

r   c                  C   s   d} t d d d d d d ddd}d d d d d ddd}t| | t| | ttjjj	| }|j
}t| t| ||fD ]*}d|_td |_td |_tjjj|_qd S )	Nr   zrand_like(Tensor self, *, ScalarType? dtype=None, Layout? layout=None, Device? device=None, bool? pin_memory=None, MemoryFormat? memory_format=None) -> Tensorrh   r   r   r   r   c                S   s"   t | j}tjj| | j|||dS r   )r
   re   r   r   r   )selfrh   r   r   r   r   r   rV   rV   rW   _meta_rand_likee  s    	
z+register_rand_like.<locals>._meta_rand_likec                S   s   t j| |||||dS )Nr   )r   r   )r   rh   r   r   r   r   rV   rV   rW   r   w  s    	z&register_rand_like.<locals>._prim_implzComputes rand_liker   )r   r   r   r   r   r   rV   rV   rW   register_rand_like]  s4    

r   c                     s   d} t d t |  dd  ddddd	}dddd
d}t| | t| | tjjjj	}|j
fdd}td| tdtjddddfdd ddd fdd}t| | |fD ]*}d|_td |_td |_tjjj|_qdS )zTThis function is used to register the var_mean function in torch.ops.nvprims module.zvar_mean.mainz7var_mean(Tensor inp, bool unbiased) -> (Tensor, Tensor)z`(Tensor inp, int[1]? dim=None, bool? unbiased=None, bool keepdim=False, *, int? correction=None)z -> (Tensor, Tensor)NFr   c          
         s   t jjrt jj}nj}t jj |d}t jj jd}|r fddtjD } fddtjD }	t j	j
j|||	}t j	j
j|||	}||fS )N)r~   c                    s"   g | ]}| vrj | nd qS rp   re   .0idiminprV   rW   
<listcomp>  s   z=register_var_mean.<locals>._meta_var_mean.<locals>.<listcomp>c                    s   g | ]}| vr|qS rV   rV   r   r   rV   rW   r         )r   r   r   rh   Zcorresponding_real_dtyper   Z_reduction_metarangendimr   rY   r   rH   )
r   r   r   r   r   r~   rL   r   output_shapebroadcast_dimsrV   r   rW   _meta_var_mean  s"    

z)register_var_mean.<locals>._meta_var_meanc                S   s    t j||}t j| |||dS )N)r   r   )r   r   set_correctionr   )r   r   r   r   r   rV   rV   rW   r     s    z%register_var_mean.<locals>._prim_implc                    s    | d |dS )N)r   r   rV   )r   r   r   rV   rW   _unbiased_overload_impl  s    z2register_var_mean.<locals>._unbiased_overload_implr   )rd   )Ztype_promoting_argsr   c          
         s   t j||}dksg kr"d t j jt j jrHtd |d}|rć fddt j	D }fddt j	D }|\}}	t j
jj|||}t j
jj|	||}	||	f}|S )NrV   r   r   c                    s"   g | ]}|vr j | nd qS r   r   r   rd   r   rV   rW   r     r   z<register_var_mean.<locals>._var_mean_ref.<locals>.<listcomp>c                    s   g | ]}| vr|qS rV   rV   r   r   rV   rW   r     r   )r   r   r   Zreduction_dimsre   r   rh   r   r   r   r   rY   r   rH   )
rd   r   r   r   r   r   r   r   rL   r   r   r   rW   _var_mean_ref  s&    

z(register_var_mean.<locals>._var_mean_refc                   sL   ddl m} | & t | ||||dW  d    S 1 s>0    Y  d S )Nr   r   r   r   )rd   r   r   r   r   r   )r   rV   rW   _var_mean_autograd  s
    
z-register_var_mean.<locals>._var_mean_autogradz]Computes the variance and mean of x over the list of dimensions specified in the dim argument)NNF)NNF)NNF)NNF)r   r   r   r   r   r   r   rY   r   r   mainnvprim_implicit_implr   r   ZCOMPLEX_TO_FLOATr   r   rO   r   r   r   r   r   r   r   )r   r   r   r   r   r   r   rV   )r   r   rW   register_var_mean  s<    
 

r   c                 C   s
   |  |S rf   Zreshaperd   Zoriginal_shaperz   rV   rV   rW   _nvprims_view_impl_aten  s    r   c                  C   s   d} t d t d dd }t| | tjjjj}|j	}dd }t
d| t| t| ||fD ]0}d	|_td |_td |_tjjj|_t|_qjd
S )zMThis function is used to register the view function in torch.ops.view module.ry   zCview(Tensor inp, SymInt[] original_shape, SymInt[] shape) -> Tensorz0view.shape(Tensor inp, SymInt[] shape) -> Tensorc                 S   s
   |  |S rf   r   r   rV   rV   rW   r     s    z!register_view.<locals>._prim_implc                 S   s6   t | jt |kr tjj| S tjjj| | j|S rf   )listre   r   rY   r   rG   ry   r   )rd   re   rV   rV   rW   %_view_no_original_shape_overload_impl  s    z<register_view.<locals>._view_no_original_shape_overload_implz
view.shapezMCreates a tensor with the specified shape containing a copy of the data in a.N)r   r   r   r   r   r   rY   r   ry   r   r   r   r   r   rO   r   r   r   r   r   ZVIEWr   r   	impl_aten)r   r   r   r   r   r   rV   rV   rW   register_view  s     



r   c                  C   s   t   t  t  t  t  tD ]} ttjj	j
| }t|j t| |j t| |j ttjj	j| }|j}t| t| ||fD ]4}|j|_t|  |_t| d|_|j|_|j|_qq"dS )zARegisters all nvFuser primitives in the torch.ops.nvprims module.FN)r   r   r   r   r   nvprim_namesrS   r   r   rY   Zprimsr   r   Zschemar   r   Z	prim_implr   Zprim_meta_implr   r   r   r   r   rO   r   r   getr   r   r   )r   Z	main_primr   r   r   rV   rV   rW   register_nvprims&  s&    
r   )NF)Ctypingr   r   r   r   r   Ztorch._prims_commonr   r   r   r   r   r	   r
   r   r   r   Ztorch._prims_common.wrappersr   r   r   Znvprim_namespaceZlibraryLibraryr   r   r   r   r   r   rO   str__annotations__Z_nvfuser_unary_opsrX   rP   execZ_nvfuser_binary_opsZ_nvfuser_ternary_opsrb   rg   rh   rk   rn   rt   rx   r{   r   intr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   rV   rV   rV   rW   <module>   s  (@%
	
  
		EJO;g#