
    i                         d dl mZmZmZ d dlmZ d dlmZ d dlm	Z	m
Z
 d dlmZ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y)    )AnyCallableOptional)Literal)PermutationInvariantTraining)#ScaleInvariantSignalDistortionRatioSignalDistortionRatio)ScaleInvariantSignalNoiseRatioSignalNoiseRatio)_deprecated_root_import_classc                   J     e Zd ZdZ	 	 ddeded   ded   dedd	f
 fd
Z xZS )_PermutationInvariantTraininga  Wrapper for deprecated import.

    >>> import torch
    >>> from torchmetrics.functional import scale_invariant_signal_noise_ratio
    >>> preds = torch.randn(3, 2, 5) # [batch, spk, time]
    >>> target = torch.randn(3, 2, 5) # [batch, spk, time]
    >>> pit = _PermutationInvariantTraining(scale_invariant_signal_noise_ratio,
    ...     mode="speaker-wise", eval_func="max")
    >>> pit(preds, target)
    tensor(-2.1065)

    metric_funcmode)speaker-wisezpermutation-wise	eval_func)maxminkwargsreturnNc                 D    t        dd       t        |   d|||d| y )Nr   audio)r   r   r    r   super__init__)selfr   r   r   r   	__class__s        s/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchmetrics/audio/_deprecated.pyr   z&_PermutationInvariantTraining.__init__   s*     	&&DgN[[ty[TZ[    )r   r   )	__name__
__module____qualname____doc__r   r   r   r   __classcell__r   s   @r   r   r      s]      =K+0	\\ 89\ <(	\
 \ 
\ \r    r   c                   4     e Zd ZdZ	 ddededdf fdZ xZS )$_ScaleInvariantSignalDistortionRatioa  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> target = tensor([3.0, -0.5, 2.0, 7.0])
    >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
    >>> si_sdr = _ScaleInvariantSignalDistortionRatio()
    >>> si_sdr(preds, target)
    tensor(18.4030)

    	zero_meanr   r   Nc                 @    t        dd       t        |   dd|i| y )Nr   r   r)   r   r   r   r)   r   r   s      r   r   z-_ScaleInvariantSignalDistortionRatio.__init__0   s%    
 	&&KWU7977r    Fr!   r"   r#   r$   boolr   r   r%   r&   s   @r   r(   r(   $   3    	  88 8 
	8 8r    r(   c                   ,     e Zd ZdZdeddf fdZ xZS )_ScaleInvariantSignalNoiseRatioa  Wrapper for deprecated import.

    >>> from torch import tensor
    >>> target = tensor([3.0, -0.5, 2.0, 7.0])
    >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
    >>> si_snr = _ScaleInvariantSignalNoiseRatio()
    >>> si_snr(preds, target)
    tensor(15.0918)

    r   r   Nc                 <    t        dd       t        |   di | y )Nr
   r   r   r   )r   r   r   s     r   r   z(_ScaleInvariantSignalNoiseRatio.__init__E   s      	&&FP"6"r    )r!   r"   r#   r$   r   r   r%   r&   s   @r   r1   r1   9   s$    	## 
# #r    r1   c                   R     e Zd ZdZ	 	 	 	 d
dee   dededee   deddf fd	Z	 xZ
S )_SignalDistortionRatioa?  Wrapper for deprecated import.

    >>> import torch
    >>> preds = torch.randn(8000)
    >>> target = torch.randn(8000)
    >>> sdr = _SignalDistortionRatio()
    >>> sdr(preds, target)
    tensor(-11.9930)
    >>> # use with pit
    >>> from torchmetrics.functional import signal_distortion_ratio
    >>> preds = torch.randn(4, 2, 8000)  # [batch, spk, time]
    >>> target = torch.randn(4, 2, 8000)
    >>> pit = _PermutationInvariantTraining(signal_distortion_ratio,
    ...     mode="speaker-wise", eval_func="max")
    >>> pit(preds, target)
    tensor(-11.7277)

    Nuse_cg_iterfilter_lengthr)   	load_diagr   r   c                 F    t        dd       t        |   d||||d| y )Nr	   r   )r5   r6   r)   r7   r   r   )r   r5   r6   r)   r7   r   r   s         r   r   z_SignalDistortionRatio.__init__a   s4     	&&=wG 	
#=Iaj	
nt	
r    )Ni   FN)r!   r"   r#   r$   r   intr.   floatr   r   r%   r&   s   @r   r4   r4   M   sb    * &* %)
c]
 
 	

 E?
 
 

 
r    r4   c                   4     e Zd ZdZ	 ddededdf fdZ xZS )_SignalNoiseRatiozWrapper for deprecated import.

    >>> from torch import tensor
    >>> target = tensor([3.0, -0.5, 2.0, 7.0])
    >>> preds = tensor([2.5, 0.0, 2.0, 8.0])
    >>> snr = _SignalNoiseRatio()
    >>> snr(preds, target)
    tensor(16.1805)

    r)   r   r   Nc                 @    t        dd       t        |   dd|i| y )Nr   r   r)   r   r   r+   s      r   r   z_SignalNoiseRatio.__init__{   s%    
 	&&8'B7977r    r,   r-   r&   s   @r   r<   r<   o   r/   r    r<   N)typingr   r   r   typing_extensionsr   torchmetrics.audio.pitr   torchmetrics.audio.sdrr   r	   torchmetrics.audio.snrr
   r   torchmetrics.utilities.printsr   r   r(   r1   r4   r<   r   r    r   <module>rD      s^    * * % ? ] S G\$@ \28+N 8*#&D #(
2 
D8( 8r    