a
    dB                  	   @   s  d dl Z d dlmZmZmZmZmZmZmZm	Z	 d dl
Z
d dlmZ d dlmZ d dlmZ d dlmZmZmZmZ d dlmZmZ g dZe	ded ed ed	 f Zh d
Ze
jjjZeee e!edddZ"e
j#e!e
j#dddZ$daee!edddZ%eee df ee df edddZ&e'eee  e ee!edddZ(e'eee  e ee!e!edddZ)e'eee  e ee!eddd Z*eej+e dbeee  e eed"d#d$Z,eej-e dceee  e eed"d%d&Z.eej/e ddeee  e eed"d'd(Z0eej1e deeee  e eed"d)d*Z2eej3e dfeee  e eed"d+d,Z4eej5e dgeee  e eed"d-d.Z6G d/d0 d0eZ7eee ee e7d1d2d3Z8ee  e d4d5d6Z9e'eee df ee df ee!ed7d8d9Z:eej;e dheee ee eed:d;d<Z<eej=e dieee ee eed:d=d>Z>eej?e djeee ee eed:d?d@Z@eejAe dkeee ee eed:dAdBZBG dCdD dDeZCe'eee ee eCdEdFdGZDeejEe dleee ee eed:dHdIZFeejGe dmeee ee eed:dJdKZHeejIe dneee ee eed:dMdNZJeejKe doeee ee eed:dOdPZLeejMe dpeee ee eed:dQdRZNeejOe dqeee ee eed:dSdTZPeejQe dreee ee eed:dUdVZReejSe dseee ee eed:dWdXZTee eee  dYdZd[ZUeejVdteee ed\d]d^ZWeejXdueee ed\d_d`ZYdS )v    N)IterableListLiteral
NamedTupleOptionalSequenceTupleUnion)register_decomposition)checkDimsType	ShapeTypeTensorLikeType)_maybe_convert_to_dtypeout_wrapper)fftfft2fftnhffthfft2hfftnrfftrfft2rfftnifftifft2ifftnihfftihfft2ihfftnirfftirfft2irfftnfftshift	ifftshiftforwardbackwardortho>   r'   r%   Nr&   )xnormsignal_numelr%   returnc                    sf   t  tv  fdd  dkr0| dt|  S | rF du pP dkpP|oP dk}|rb| d|  S | S )z3Apply normalization to the un-normalized FFT resultc                      s
   d  S )NzInvalid normalization mode:  r,   r)   r,   X/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/torch/_refs/fft.py<lambda>.       z_apply_norm.<locals>.<lambda>r'      Nr&   r%   )r   _NORM_VALUESmathsqrt)r(   r)   r*   r%   	normalizer,   r-   r.   _apply_norm*   s    
r6   )dtyperequire_complexr+   c                 C   s*   | j r
| S | jst } |r&t| } | S )z@Helper to promote a dtype to one supported by the FFT primitives)
is_complexZis_floating_pointtorchZget_default_dtypeutilsZcorresponding_complex_dtype)r7   r8   r,   r,   r.   _promote_type_fft9   s    
r<   F)tr8   r+   c                 C   s   | j }t||}t| |S )zEHelper to promote a tensor to a dtype supported by the FFT primitives)r7   r<   r   )r=   r8   Zcur_typenew_typer,   r,   r.   _maybe_promote_tensor_fftH   s    
r?   .)r(   dimssizesr+   c                 C   s   t |t |ksJ d}| j}dgt | d }tt |D ]}|| dkrNq<|||  || k rd}t |d||   d }|| |||   ||< |||  || kr<| || d|| } q<|rt| |S | S )z
    Fixes the shape of x such that x.size(dims[i]) == sizes[i],
    either by zero-padding, or by slicing x starting from 0.
    Fr      Tr1   )lenshaperangeZnarrowr:   Zconstant_pad_nd)r(   r@   rA   Z	must_copyZx_sizesZ
pad_amountiZpad_idxr,   r,   r.   _resize_fft_inputQ   s    rH   )	func_nameinputndimr)   r%   r+   c           	         s   t |dd}tj|j|ddf} dur, nd|j| d  }t|dk fdd	  durtt|||d d fd
}|rt|}t	j
|||d}t||||dS )zBCommon code for performing any complex to real FFT (irfft or hfft)Tr8   FZwrap_scalarNrB   r1   c                      s   d  dS NzInvalid number of data points (z) specifiedr,   r,   rK   r,   r.   r/   w   r0   z_fft_c2r.<locals>.<lambda>)r@   rA   rL   last_dim_sizer)   r*   r%   )r?   r;   canonicalize_dimndimrE   r   rH   r:   conjprimsfft_c2rr6   )	rI   rJ   rK   rL   r)   r%   r@   rR   outputr,   rP   r.   _fft_c2rk   s    	
rZ   )rI   rJ   rK   rL   r)   r%   onesidedr+   c           	         s   t jj  fdd ttjj|ddf}|durLt||ftj	||d}t
||j| |}|rx|S t|S )zBCommon code for performing any real to complex FFT (rfft or ihfft)c                      s     dj  S )Nz0 expects a floating point input tensor, but got r7   r,   rI   rJ   r,   r.   r/      r0   z_fft_r2c.<locals>.<lambda>FrN   NrL   r[   )r   r7   r9   r?   r;   rT   rU   rH   rW   fft_r2cr6   rE   r:   rV   )	rI   rJ   rK   rL   r)   r%   r[   r@   retr,   r]   r.   _fft_r2c   s    
ra   c                    sf   t jj fdd tjj|ddf}|durBt||ftj||d}t	||j
| |S )zCCommon code for performing any complex to complex FFT (fft or ifft)c                      s     dj  S Nz) expects a complex input tensor, but got r\   r,   r]   r,   r.   r/      r0   z_fft_c2c.<locals>.<lambda>FrN   NrL   r%   )r   r7   r9   r;   rT   rU   rH   rW   fft_c2cr6   rE   )rI   rJ   rK   rL   r)   r%   r@   r`   r,   r]   r.   _fft_c2c   s    	re   rC   )rJ   rK   rL   r)   r+   c              	   C   s6   | j jrtd| |||ddS td| |||dddS d S )Nr   Tr%   Fr%   r[   r7   r9   re   ra   rJ   rK   rL   r)   r,   r,   r.   r      s    r   c              	   C   s6   | j jrtd| |||ddS td| |||dddS d S )Nr   Frf   rg   rh   ri   r,   r,   r.   r      s    r   c              	   C   s   t d| |||dddS )Nr   Trg   ra   ri   r,   r,   r.   r      s    r   c                 C   s   t d| |||ddS )Nr    Frf   rZ   ri   r,   r,   r.   r       s    r    c                 C   s   t d| |||ddS )Nr   Trf   rk   ri   r,   r,   r.   r      s    r   c              	   C   s   t d| |||dddS )Nr   FTrg   rj   ri   r,   r,   r.   r      s    r   c                   @   s.   e Zd ZU eedf ed< eedf ed< dS )_ShapeAndDims.rE   r@   N__name__
__module____qualname__r   int__annotations__r,   r,   r,   r.   rl      s   
rl   )rJ   rE   rL   r+   c                    sH  | j  | j|durRt|ts$|f}tj |dd}ttt|t|kdd  |durt|tsj|f}t|du pt|t|kdd  t|t k fdd |du rt	t
   }t	fdd	t||D }n6|du rt	t
 }t	}nt	fd
d	|D }|D ]tdkfdd qt||dS )zTConvert the shape and dim arguments into a canonical form where neither are optionalNFrN   c                   S   s   dS )NzFFT dims must be uniquer,   r,   r,   r,   r.   r/     r0   z6_canonicalize_fft_shape_and_dim_args.<locals>.<lambda>c                   S   s   dS )Nz=When given, dim and shape arguments must have the same lengthr,   r,   r,   r,   r.   r/     r0   c                      s   d d  dS )NzGot shape with z" values but input tensor only has z dimensions.r,   r,   )	input_dimtransform_ndimr,   r.   r/     s   c                 3   s&   | ]\}}|d kr|n | V  qdS )rC   Nr,   ).0sdinput_sizesr,   r.   	<genexpr>$  s   z7_canonicalize_fft_shape_and_dim_args.<locals>.<genexpr>c                 3   s   | ]} | V  qd S Nr,   ru   rw   rx   r,   r.   rz   -  r0   r   c                      s   d  dS rO   r,   r,   rP   r,   r.   r/   0  r0   )rE   r@   )rU   rE   
isinstancer   r;   Zcanonicalize_dimsr   rD   settuplerF   ziprl   )rJ   rE   rL   Zret_dimsZ	ret_shaper,   )rs   ry   rK   rt   r.   $_canonicalize_fft_shape_and_dim_args   s>    




r   )xsr+   c                 C   s   d}| D ]}||9 }q|S )zCompute product of a listr1   r,   )r   prodr(   r,   r,   r.   _prod5  s    
r   )function_namerJ   rE   rL   r)   r%   r+   c                    sH   t jj fdd t||}tj|||d}t||t||dS )zECommon code for n-dimensional complex to complex FFTs (fftn or ifftn)c                      s     dj  S rb   r\   r,   r   rJ   r,   r.   r/   H  s   z_fftn_c2c.<locals>.<lambda>rc   rS   )r   r7   r9   rH   rW   rd   r6   r   )r   rJ   rE   rL   r)   r%   r(   rY   r,   r   r.   	_fftn_c2c=  s    	r   )rJ   rv   rL   r)   r+   c                 C   s0   t | ||\}}t| dd}td||||ddS )NTrM   r   rf   r   r?   r   rJ   rv   rL   r)   rE   r(   r,   r,   r.   r   P  s    r   c                 C   s0   t | ||\}}t| dd}td||||ddS )NTrM   r   Frf   r   r   r,   r,   r.   r   ]  s    r   c                    sd   t  jj  fdd t ||\}}t dd t || tj |dd}t||t	|ddS )Nc                      s   d j  S )Nz2rfftn expects a real-valued input tensor, but got r\   r,   rJ   r,   r.   r/   t  r0   zrfftn.<locals>.<lambda>FrM   Tr^   rS   )
r   r7   r9   r   r?   rH   rW   r_   r6   r   )rJ   rv   rL   r)   rE   outr,   r   r.   r   j  s    
r   c                    s   t  jj  fdd t ||\}}t t|dkdd  t dd t || tj |dd  dd	}t|d
krt	|||d dd}t
|S t|}tj||d d dd}t	||t|ddS )Nc                      s   d j  S )Nz3ihfftn expects a real-valued input tensor, but got r\   r,   r   r,   r.   r/     r0   zihfftn.<locals>.<lambda>r   c                   S   s   dS )Nz'ihfftn must transform at least one axisr,   r,   r,   r,   r.   r/     r0   FrM   rC   Tr^   r1   rS   rc   )r   r7   r9   r   rD   r?   rH   rW   r_   r6   rV   conj_physicalrd   r   )rJ   rv   rL   r)   rE   tmpr,   r   r.   r   }  s    


r   c                   @   s6   e Zd ZU eedf ed< eedf ed< eed< dS )_CanonicalizeC2rReturn.rE   rL   rR   Nrm   r,   r,   r,   r.   r     s   
r   )fnamerJ   rv   rL   r+   c                    s   t |||\}}tt|dk fdd |du s>|d dkrVd|j|d  d  n|d tdkfdd t|}d d |d< tt||d	S )
zCanonicalize shape and dim arguments for n-dimensional c2r transforms,
    as well as calculating the last_dim_size which is shape[dim[-1]] for the outputr   c                      s
     dS )Nz! must transform at least one axisr,   r,   )r   r,   r.   r/     r0   z:_canonicalize_fft_c2r_shape_and_dim_args.<locals>.<lambda>NrC   rB   r1   c                      s   d  dS rO   r,   r,   )rR   r,   r.   r/     r0   )rE   rL   rR   )r   r   rD   rE   listr   r   )r   rJ   rv   rL   rE   Z
shape_listr,   )r   rR   r.   (_canonicalize_fft_c2r_shape_and_dim_args  s    

r   c                    s^   t d| ||\}}}t| dd} t| ||} tj| ||d t |t fdd|D ddS )	Nr"   TrM   rQ   c                 3   s   | ]} j | V  qd S r{   rE   r|   r   r,   r.   rz     r0   zirfftn.<locals>.<genexpr>Frf   )r   r?   rH   rW   rX   r6   r   )rJ   rv   rL   r)   rE   rR   r,   r   r.   r"     s    
r"   c                 C   s   t d| ||\}}}t| dd} t| ||} t|dkrPtj| |d d ddn| }t||t|d d dd}t|}tj	||dd  |d}t|||ddS )	Nr   TrM   r1   rC   rc   rf   rQ   )
r   r?   rH   rD   rW   rd   r6   r   r   rX   )rJ   rv   rL   r)   rE   rR   r   r   r,   r,   r.   r     s    
(
r   rC   c                 C   s   t jj| |||dS N)rv   rL   r)   )r:   r   r   rJ   rv   rL   r)   r,   r,   r.   r     s    r   c                 C   s   t jj| |||dS r   )r:   r   r   r   r,   r,   r.   r     s    r   c                 C   s   t jj| |||dS r   )r:   r   r   r   r,   r,   r.   r     s    r   c                 C   s   t jj| |||dS r   )r:   r   r"   r   r,   r,   r.   r!     s    r!   c                 C   s   t jj| |||dS r   )r:   r   r   r   r,   r,   r.   r     s    r   c                 C   s   t jj| |||dS r   )r:   r   r   r   r,   r,   r.   r     s    r   )rL   r(   r+   c                 C   s2   | du rt t|jS t| ts&| gS t | S dS )zIConvert Optional[DimsType] to a simple list, defaulting to all dimensionsN)r   rF   rU   r}   r   )rL   r(   r,   r,   r.   _default_alldims#  s
    
r   )rJ   rL   r+   c                    s*   t | } fdd|D }t ||S )Nc                    s   g | ]} j | d  qS )rB   r   r|   r   r,   r.   
<listcomp>0  r0   zfftshift.<locals>.<listcomp>r   r:   ZrollrJ   rL   r@   shiftr,   r   r.   r#   -  s    
r#   c                    s*   t | } fdd|D }t ||S )Nc                    s   g | ]} j | d  d qS )r1   rB   r   r|   r   r,   r.   r   7  r0   zifftshift.<locals>.<listcomp>r   r   r,   r   r.   r$   4  s    
r$   )F)NrC   N)NrC   N)NrC   N)NrC   N)NrC   N)NrC   N)NNN)NNN)NNN)NNN)NNN)NNN)Nr   N)Nr   N)Nr   N)Nr   N)Nr   N)Nr   N)N)N)Zr3   typingr   r   r   r   r   r   r   r	   r:   Ztorch._primsZ_primsrW   Ztorch._prims_commonZ_prims_commonr;   Ztorch._decompr
   r   r   r   r   Ztorch._prims_common.wrappersr   r   __all__ZNormTyper2   Z_opsopsZatenrq   boolr6   r7   r<   r?   rH   strrZ   ra   re   Zfft_fftr   Zfft_ifftr   Zfft_rfftr   Z	fft_irfftr    Zfft_hfftr   Z	fft_ihfftr   rl   r   r   r   Zfft_fftnr   Z	fft_ifftnr   Z	fft_rfftnr   Z
fft_ihfftnr   r   r   Z
fft_irfftnr"   Z	fft_hfftnr   Zfft_fft2r   Z	fft_ifft2r   Z	fft_rfft2r   Z
fft_irfft2r!   Z	fft_hfft2r   Z
fft_ihfft2r   r   Zfft_fftshiftr#   Zfft_ifftshiftr$   r,   r,   r,   r.   <module>   sB  (
 
         	   	   	   	6	

                     	   	   	   	   	   	
