
    i8!                        d dl mZ d dlmZmZmZ d dlmZ d dlm	Z	 d dl
mZ d dlmZ d dlmZ d dlmZ d d	lmZ d d
lmZmZ d dlmZ d dlmZ d dlmZ  G d de      Z G d de      Z G d de      Z G d de      Z  G d de      Z! G d de      Z" G d de	      Z# G d de      Z$ G d de      Z% G d  d!e      Z&y")#    )Sequence)AnyOptionalUnion)Literal)SpectralDistortionIndex))ErrorRelativeGlobalDimensionlessSynthesis)PeakSignalNoiseRatio)RelativeAverageSpectralError)&RootMeanSquaredErrorUsingSlidingWindow)SpectralAngleMapper)*MultiScaleStructuralSimilarityIndexMeasure StructuralSimilarityIndexMeasure)TotalVariation)UniversalImageQualityIndex)_deprecated_root_import_classc            	       @     e Zd ZdZ	 	 d	deded   deddf fdZ xZS )
*_ErrorRelativeGlobalDimensionlessSynthesiszWrapper for deprecated import.

    >>> from torch import rand
    >>> preds = rand([16, 1, 16, 16])
    >>> target = preds * 0.75
    >>> ergas = _ErrorRelativeGlobalDimensionlessSynthesis()
    >>> ergas(preds, target).round()
    tensor(10.)

    ratio	reductionelementwise_meansumnoneNkwargsreturnNc                 B    t        dd       t        |   d||d| y )Nr	   image)r   r    r   super__init__)selfr   r   r   	__class__s       s/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/image/_deprecated.pyr"   z3_ErrorRelativeGlobalDimensionlessSynthesis.__init__   s(     	&&QSZ[Du	DVD    )   r   )	__name__
__module____qualname____doc__floatr   r   r"   __classcell__r$   s   @r%   r   r      sL    	 FXEE BCE 	E
 
E Er&   r   c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 ddedeeee   f   deeee   f   de	d   de
eeeeef   f      d	ed
edeedf   de	d   deddf fdZ xZS )+_MultiScaleStructuralSimilarityIndexMeasurea	  Wrapper for deprecated import.

    >>> from torch import rand
    >>> preds = rand([3, 3, 256, 256])
    >>> target = preds * 0.75
    >>> ms_ssim = _MultiScaleStructuralSimilarityIndexMeasure(data_range=1.0)
    >>> ms_ssim(preds, target)
    tensor(0.9628)

    Ngaussian_kernelkernel_sizesigmar   r   
data_rangek1k2betas.	normalize)relusimpleNr   r   c
                 P    t        dd       t        |   d|||||||||	d	|
 y )Nr   r   )	r1   r2   r3   r   r4   r5   r6   r7   r8   r   r    )r#   r1   r2   r3   r   r4   r5   r6   r7   r8   r   r$   s              r%   r"   z4_MultiScaleStructuralSimilarityIndexMeasure.__init__4   sH     	&&RT[\ 	
+#!	
 	
r&   )	T         ?r   N{Gz?Q?)gǺ?g48EG?ga4?g??g9EGr?r9   )r(   r)   r*   r+   boolr   intr   r,   r   r   tupler   r"   r-   r.   s   @r%   r0   r0   (   s    	 !%13/2FXBF#K5;

 3-.
 UHUO+,	

 BC
 U5%u*=#=>?
 
 
 UCZ 
 12
 
 

 
r&   r0   c                   z     e Zd ZdZ	 	 	 	 ddeeeeef   f   deded   deee	ee	df   f      d	e
d
df fdZ xZS )_PeakSignalNoiseRatiozWrapper for deprecated import.

    >>> from torch import tensor
    >>> psnr = _PeakSignalNoiseRatio()
    >>> preds = tensor([[0.0, 1.0], [2.0, 3.0]])
    >>> target = tensor([[3.0, 2.0], [1.0, 0.0]])
    >>> psnr(preds, target)
    tensor(2.5527)

    Nr4   baser   r   dim.r   r   c                 F    t        dd       t        |   d||||d| y )Nr
   r   )r4   rE   r   rF   r   r    )r#   r4   rE   r   rF   r   r$   s         r%   r"   z_PeakSignalNoiseRatio.__init__\   s-     	&&<gFbJTYTWb[abr&   )g      @g      $@r   N)r(   r)   r*   r+   r   r,   rB   r   r   rA   r   r"   r-   r.   s   @r%   rD   rD   P   s    	 9<FX59	c%ue|!445	c 	c BC		c
 eCsCx012	c 	c 
	c 	cr&   rD   c                   >     e Zd ZdZ	 ddedeeef   ddf fdZ xZ	S )_RelativeAverageSpectralErrorzWrapper for deprecated import.

    >>> from torch import rand
    >>> preds = rand(4, 3, 16, 16)
    >>> target = rand(4, 3, 16, 16)
    >>> rase = _RelativeAverageSpectralError()
    >>> rase(preds, target)
    tensor(5326.40...)

    window_sizer   r   Nc                 @    t        dd       t        |   dd|i| y )Nr   r   rJ   r   r    r#   rJ   r   r$   s      r%   r"   z&_RelativeAverageSpectralError.__init__t   s%    
 	&&DgN;[;F;r&      
r(   r)   r*   r+   rA   dictstrr   r"   r-   r.   s   @r%   rI   rI   h   ;    	 << sCx.< 
	< <r&   rI   c                   >     e Zd ZdZ	 ddedeeef   ddf fdZ xZ	S )'_RootMeanSquaredErrorUsingSlidingWindowzWrapper for deprecated import.

    >>> from torch import rand
    >>> preds = rand(4, 3, 16, 16)
    >>> target = rand(4, 3, 16, 16)
    >>> rmse_sw = RootMeanSquaredErrorUsingSlidingWindow()
    >>> rmse_sw(preds, target)
    tensor(0.4158)

    rJ   r   r   Nc                 @    t        dd       t        |   dd|i| y )Nr   r   rJ   r   r    rL   s      r%   r"   z0_RootMeanSquaredErrorUsingSlidingWindow.__init__   s&    
 	&&NPWX;[;F;r&   rM   rO   r.   s   @r%   rT   rT   }   rR   r&   rT   c                   :     e Zd ZdZ	 dded   deddf fdZ xZS )	_SpectralAngleMapperzWrapper for deprecated import.

    >>> from torch import rand
    >>> preds = rand([16, 3, 16, 16])
    >>> target = rand([16, 3, 16, 16])
    >>> sam = _SpectralAngleMapper()
    >>> sam(preds, target)
    tensor(0.5914)

    r   r   r   r   r   r   Nc                 @    t        dd       t        |   dd|i| y )Nr   r   r   r   r    r#   r   r   r$   s      r%   r"   z_SpectralAngleMapper.__init__   s%    
 	&&;WE7977r&   )r   r(   r)   r*   r+   r   r   r"   r-   r.   s   @r%   rW   rW      s;    	 AS8<=8 8 
	8 8r&   rW   c            	       >     e Zd ZdZ	 d	deded   deddf fdZ xZS )
_SpectralDistortionIndexzWrapper for deprecated import.

    >>> from torch import rand
    >>> preds = rand([16, 3, 16, 16])
    >>> target = rand([16, 3, 16, 16])
    >>> sdi = _SpectralDistortionIndex()
    >>> sdi(preds, target)
    tensor(0.0234)

    pr   rX   r   r   Nc                 B    t        dd       t        |   d||d| y )Nr   r   )r^   r   r   r    )r#   r^   r   r   r$   s       r%   r"   z!_SpectralDistortionIndex.__init__   s'     	&&?I<1	<V<r&   )   r   )	r(   r)   r*   r+   rA   r   r   r"   r-   r.   s   @r%   r]   r]      s?    	 Se==%,-N%O=ps=	= =r&   r]   c                        e Zd ZdZ	 	 	 	 	 	 	 	 	 ddedeeee   f   deeee   f   de	d   de
eeeeef   f      d	ed
ededededdf fdZ xZS )!_StructuralSimilarityIndexMeasurezWrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([3, 3, 256, 256])
    >>> target = preds * 0.75
    >>> ssim = _StructuralSimilarityIndexMeasure(data_range=1.0)
    >>> ssim(preds, target)
    tensor(0.9219)

    Nr1   r3   r2   r   r   r4   r5   r6   return_full_imagereturn_contrast_sensitivityr   r   c
                 P    t        dd       t        |   d|||||||||	d	|
 y )Nr   r   )	r1   r3   r2   r   r4   r5   r6   rc   rd   r   r    )r#   r1   r3   r2   r   r4   r5   r6   rc   rd   r   r$   s              r%   r"   z*_StructuralSimilarityIndexMeasure.__init__   sG     	&&H'R 	
+#!/(C	
 	
r&   )	Tr=   r<   r   Nr>   r?   FF)r(   r)   r*   r+   r@   r   r,   r   rA   r   r   rB   r   r"   r-   r.   s   @r%   rb   rb      s    	 !%/213FXBF"',1

 UHUO+,
 3-.	

 BC
 U5%u*=#=>?
 
 
  
 &*
 
 

 
r&   rb   c                   8     e Zd ZdZdded   deddf fdZ xZS )	_TotalVariationzWrapper for deprecated import.

    >>> from torch import rand
    >>> tv = _TotalVariation()
    >>> img = rand(5, 3, 28, 28)
    >>> tv(img)
    tensor(7546.8018)

    r   )meanr   r   Nr   r   Nc                 @    t        dd       t        |   dd|i| y )Nr   r   r   r   r    rZ   s      r%   r"   z_TotalVariation.__init__   s#    %&6@7977r&   )r   r[   r.   s   @r%   rg   rg      s/    8'*E"F 8Z] 8bf 8 8r&   rg   c                   R     e Zd ZdZ	 	 	 d
dee   dee   ded   deddf
 fd	Z	 xZ
S )_UniversalImageQualityIndexzWrapper for deprecated import.

    >>> import torch
    >>> preds = torch.rand([16, 1, 16, 16])
    >>> target = preds * 0.75
    >>> uqi = _UniversalImageQualityIndex()
    >>> uqi(preds, target)
    tensor(0.9216)

    r2   r3   r   r   r   r   Nc                 D    t        dd       t        |   d|||d| y )Nr   r   )r2   r3   r   r   r    )r#   r2   r3   r   r   r$   s        r%   r"   z$_UniversalImageQualityIndex.__init__   s*     	&&BGL][]V\]r&   ))r<   r<   )r=   r=   r   )r(   r)   r*   r+   r   rA   r,   r   r   r"   r-   r.   s   @r%   rk   rk      sb    	 &.!+FX	^c]^ ^ BC	^
 ^ 
^ ^r&   rk   N)'collections.abcr   typingr   r   r   typing_extensionsr   torchmetrics.image.d_lambdar   torchmetrics.image.ergasr	   torchmetrics.image.psnrr
   torchmetrics.image.raser   torchmetrics.image.rmse_swr   torchmetrics.image.samr   torchmetrics.image.ssimr   r   torchmetrics.image.tvr   torchmetrics.image.uqir   torchmetrics.utilities.printsr   r   r0   rD   rI   rT   rW   r]   rb   rg   rk   r   r&   r%   <module>rz      s    $ ' ' % ? N 8 @ M 6 p 0 = GE1Z E,%
2\ %
Pc0 c0<$@ <*<.T <*8. 8*=6 =&%
(H %
P8n 8 ^"< ^r&   