Ë
    óÿæiù/  ã                   ó  — U d dl Z d dlmZmZmZ d dlZd dlmZ d dlmc m	Z
 dedefd„Zd ed«      fZeeef   ed<    G d	„ d
ej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Z G d„ dej                  «      Zdedededej0                  j                  fd„Z	 d"dededededededededee   defd „Zdefd!„Zy)#é    N)ÚListÚOptionalÚTupleÚxÚreturnc                 óJ   — dddt        j                  d| z  dz   «      z   z  z   S )aŒ  The metric defined by ITU-T P.862 is often called 'PESQ score', which is defined
    for narrow-band signals and has a value range of [-0.5, 4.5] exactly. Here, we use the metric
    defined by ITU-T P.862.2, commonly known as 'wide-band PESQ' and will be referred to as "PESQ score".

    Args:
        x (float): Narrow-band PESQ score.

    Returns:
        (float): Wide-band PESQ score.
    g+‡ÙÎ÷ï?gÿÿÿÿÿÿ@é   g;pÎˆÒÞõ¿gÜ×�sF”@)ÚmathÚexp)r   s    úv/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/squim/objective.pyÚtransform_wb_pesq_ranger   	   s+   € ð �M a¬$¯(©(°7¸Q±;ÀÑ3GÓ*HÑ&HÑIÑIÐIó    ç      ð?g      @Ú	PESQRangec                   ól   ‡ — e Zd Zddeeef   ddfˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )ÚRangeSigmoidÚ	val_ranger   Nc                 óª   •— t         t        | �  «        t        |t        «      rt        |«      dk(  sJ ‚|| _        t        j                  «       | _	        y )Né   )
Úsuperr   Ú__init__Ú
isinstanceÚtupleÚlenr   ÚnnÚSigmoidÚsigmoid)Úselfr   Ú	__class__s     €r   r   zRangeSigmoid.__init__    s?   ø€ ÜŒl˜DÑ*Ô,Ü˜)¤UÔ+´°I³À!Ò0CÐCÐCØ.7ˆŒÜ*,¯*©*«,ˆ�r   r   c                 óˆ   — | j                  |«      | j                  d   | j                  d   z
  z  | j                  d   z   }|S )Nr	   r   )r   r   ©r   r   Úouts      r   ÚforwardzRangeSigmoid.forward&   s?   € Ø�l‰l˜1‹o §¡°Ñ!2°T·^±^ÀAÑ5FÑ!FÑGÈ$Ï.É.ÐYZÑJ[Ñ[ˆØˆ
r   ))ç        r   )
Ú__name__Ú
__module__Ú__qualname__r   Úfloatr   ÚtorchÚTensorr#   Ú__classcell__©r   s   @r   r   r      s:   ø„ ñ7 %¨¨u¨Ñ"5ð 7Àtõ 7ð˜Ÿ™ð ¨%¯,©,÷ r   r   c                   ój   ‡ — e Zd ZdZd	dededdfˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )
ÚEncoderzýEncoder module that transform 1D waveform to 2D representations.

    Args:
        feat_dim (int, optional): The feature dimension after Encoder module. (Default: 512)
        win_len (int, optional): kernel size in the Conv1D layer. (Default: 32)
    Úfeat_dimÚwin_lenr   Nc                 ón   •— t         t        | �  «        t        j                  d|||dz  d¬«      | _        y )Nr	   r   F)ÚstrideÚbias)r   r.   r   r   ÚConv1dÚconv1d)r   r/   r0   r   s      €r   r   zEncoder.__init__3   s-   ø€ ÜŒg�tÑ%Ô'ä—i‘i  8¨W¸WÈ¹\ÐPUÔVˆ�r   r   c                 ór   — |j                  d¬«      }t        j                  | j                  |«      «      }|S )a  Apply waveforms to convolutional layer and ReLU layer.

        Args:
            x (torch.Tensor): Input waveforms. Tensor with dimensions `(batch, time)`.

        Returns:
            (torch,Tensor): Feature Tensor with dimensions `(batch, channel, frame)`.
        r	   ©Údim)Ú	unsqueezeÚFÚrelur5   r!   s      r   r#   zEncoder.forward8   s0   € ð �k‰k˜aˆkÓ ˆÜ�f‰f�T—[‘[ Ó%Ó&ˆØˆ
r   )i   é    )
r%   r&   r'   Ú__doc__Úintr   r)   r*   r#   r+   r,   s   @r   r.   r.   +   sA   ø„ ññW ð W°Sð WÀ$õ Wð
˜Ÿ™ð ¨%¯,©,÷ r   r.   c                   ón   ‡ — e Zd Zd
dededededdf
ˆ fd„Zdej                  dej                  fd	„Z	ˆ xZ
S )Ú	SingleRNNÚrnn_typeÚ
input_sizeÚhidden_sizeÚdropoutr   Nc                 óÒ   •— t         t        | �  «        || _        || _        || _         t        t        |«      ||d|dd¬«      | _        t        j                  |dz  |«      | _
        y )Nr	   T)rD   Úbatch_firstÚbidirectionalr   )r   r@   r   rA   rB   rC   Úgetattrr   ÚrnnÚLinearÚproj)r   rA   rB   rC   rD   r   s        €r   r   zSingleRNN.__init__G   se   ø€ ÜŒi˜Ñ'Ô)à ˆŒØ$ˆŒØ&ˆÔà&;¤g¬b°(Ó&;ØØØØØØô'
ˆŒô —I‘I˜k¨A™o¨zÓ:ˆ�	r   r   c                 óP   — | j                  |«      \  }}| j                  |«      }|S ©N)rI   rK   )r   r   r"   Ú_s       r   r#   zSingleRNN.forwardY   s%   € à—‘˜!“‰ˆˆQØ�i‰i˜‹nˆØˆ
r   )r$   )r%   r&   r'   Ústrr>   r(   r   r)   r*   r#   r+   r,   s   @r   r@   r@   F   sH   ø„ ñ; ð ;°#ð ;ÀCð ;ÐRWð ;Ðbfõ ;ð$˜Ÿ™ð ¨%¯,©,÷ r   r@   c                   óL  ‡ — e Zd ZdZ	 	 	 	 	 	 	 ddededededededed	d
fˆ fd„Zdej                  d	e	ej                  ef   fd„Z
dej                  d	e	ej                  ef   fd„Zdej                  ded	ej                  fd„Zdej                  d	ej                  fd„Zˆ xZS )ÚDPRNNaÏ  *Dual-path recurrent neural networks (DPRNN)* :cite:`luo2020dual`.

    Args:
        feat_dim (int, optional): The feature dimension after Encoder module. (Default: 64)
        hidden_dim (int, optional): Hidden dimension in the RNN layer of DPRNN. (Default: 128)
        num_blocks (int, optional): Number of DPRNN layers. (Default: 6)
        rnn_type (str, optional): Type of RNN in DPRNN. Valid options are ["RNN", "LSTM", "GRU"]. (Default: "LSTM")
        d_model (int, optional): The number of expected features in the input. (Default: 256)
        chunk_size (int, optional): Chunk size of input for DPRNN. (Default: 100)
        chunk_stride (int, optional): Stride of chunk input for DPRNN. (Default: 50)
    r/   Ú
hidden_dimÚ
num_blocksrA   Úd_modelÚ
chunk_sizeÚchunk_strider   Nc                 ó$  •— t         t        | �  «        || _        t	        j
                  g «      | _        t	        j
                  g «      | _        t	        j
                  g «      | _        t	        j
                  g «      | _	        t        |«      D ]°  }| j                  j                  t        |||«      «       | j                  j                  t        |||«      «       | j                  j                  t	        j                  d|d¬«      «       | j                  j                  t	        j                  d|d¬«      «       Œ² t	        j                  t	        j                  ||d«      t	        j                   «       «      | _        || _        || _        y )Nr	   g:Œ0âŽyE>)Úeps)r   rQ   r   rS   r   Ú
ModuleListÚrow_rnnÚcol_rnnÚrow_normÚcol_normÚrangeÚappendr@   Ú	GroupNormÚ
SequentialÚConv2dÚPReLUÚconvrU   rV   )
r   r/   rR   rS   rA   rT   rU   rV   rN   r   s
            €r   r   zDPRNN.__init__m   s  ø€ ô 	Œe�TÑ#Ô%à$ˆŒä—}‘} RÓ(ˆŒÜ—}‘} RÓ(ˆŒÜŸ™ bÓ)ˆŒÜŸ™ bÓ)ˆŒÜ�zÖ"ˆAØ�L‰L×Ñ¤	¨(°H¸jÓ IÔJØ�L‰L×Ñ¤	¨(°H¸jÓ IÔJØ�M‰M× Ñ ¤§¡¨a°¸tÔ!DÔEØ�M‰M× Ñ ¤§¡¨a°¸tÔ!DÕEð	 #ô
 —M‘MÜ�I‰I�h ¨Ó+Ü�H‰H‹Jó
ˆŒ	ð %ˆŒØ(ˆÕr   r   c                 óò   — |j                   d   }| j                  | j                  || j                  z  z   | j                  z  z
  }t        j                  || j                  || j                  z   g«      }||fS )Néÿÿÿÿ)ÚshaperU   rV   r:   Úpad)r   r   Úseq_lenÚrestr"   s        r   Ú	pad_chunkzDPRNN.pad_chunk‹   sm   € à—'‘'˜"‘+ˆà�‰ $×"3Ñ"3°gÀÇÁÑ6OÑ"OÐSW×SbÑSbÑ!bÑbˆÜ�e‰e�A˜×)Ñ)¨4°$×2CÑ2CÑ+CÐDÓEˆà�DˆyÐr   c                 ó  — | j                  |«      \  }}|j                  \  }}}|d d …d d …d | j                   …f   j                  «       j	                  ||d| j
                  «      }|d d …d d …| j                  d …f   j                  «       j	                  ||d| j
                  «      }t        j                  ||gd¬«      }|j	                  ||d| j
                  «      j                  dd«      j                  «       }||fS )Nrf   é   r7   r   )	rk   rg   rV   Ú
contiguousÚviewrU   r)   ÚcatÚ	transpose)	r   r   r"   rj   Ú
batch_sizer/   ri   Ú	segments1Ú	segments2s	            r   ÚchunkingzDPRNN.chunking”   sñ   € Ø—N‘N 1Ó%‰	ˆˆTØ(+¯	©	Ñ%ˆ
�H˜gàšš1Ð2 ×!2Ñ!2Ð 2Ð2Ð2Ñ3×>Ñ>Ó@×EÑEÀjÐRZÐ\^Ð`d×`oÑ`oÓpˆ	Øšš1˜d×/Ñ/Ñ1Ð1Ñ2×=Ñ=Ó?×DÑDÀZÐQYÐ[]Ð_c×_nÑ_nÓoˆ	Ü�i‰i˜ IÐ.°AÔ6ˆØ�h‰h�z 8¨R°·±ÓA×KÑKÈAÈqÓQ×\Ñ\Ó^ˆà�DˆyÐr   rj   c                 ó:  — |j                   \  }}}}|j                  dd«      j                  «       j                  ||d| j                  dz  «      }|d d …d d …d d …d | j                  …f   j                  «       j                  ||d«      d d …d d …| j
                  d …f   }|d d …d d …d d …| j                  d …f   j                  «       j                  ||d«      d d …d d …d | j
                   …f   }||z   }|dkD  r|d d …d d …d | …f   }|j                  «       }|S )Nr   rm   rf   r   )rg   rq   rn   ro   rU   rV   )	r   r   rj   rr   r8   rN   r"   Úout1Úout2s	            r   ÚmergingzDPRNN.mergingŸ   s  € Ø !§¡Ñˆ
�C˜˜AØ�k‰k˜!˜QÓ×*Ñ*Ó,×1Ñ1°*¸cÀ2ÀtÇÁÐYZÑGZÓ[ˆØ’1’ašÐ-˜dŸo™oÐ-Ð-Ñ.×9Ñ9Ó;×@Ñ@ÀÈSÐRTÓUÒVWÒYZÐ\`×\mÑ\mÑ\oÐVoÑpˆØ’1’aš˜DŸO™OÑ-Ð-Ñ.×9Ñ9Ó;×@Ñ@ÀÈSÐRTÓUÒVWÒYZÐ\pÐ_c×_pÑ_pÐ^pÐ\pÐVpÑqˆØ�T‰kˆØ�!Š8Ø’aš˜F˜d˜U˜F�lÑ#ˆCØ�n‰nÓˆØˆ
r   c                 ó’  — | j                  |«      \  }}|j                  \  }}}}|}t        | j                  | j                  | j
                  | j                  «      D �]"  \  }}	}
}|j                  dddd«      j                  «       j                  ||z  |d«      j                  «       } ||«      }|j                  |||d«      j                  dddd«      j                  «       } |	|«      }||z   }|j                  dddd«      j                  «       j                  ||z  |d«      j                  «       } |
|«      }|j                  |||d«      j                  dddd«      j                  «       } ||«      }||z   }�Œ% | j                  |«      }| j                  ||«      }|j                  dd«      j                  «       }|S )Nr   rm   r   r	   rf   )ru   rg   ÚziprZ   r\   r[   r]   Úpermutern   ro   rd   ry   rq   )r   r   rj   rr   rN   Údim1Údim2r"   rZ   r\   r[   r]   Úrow_inÚrow_outÚcol_inÚcol_outs                   r   r#   zDPRNN.forwardª   s«  € Ø—-‘- Ó"‰ˆˆ4Ø$%§G¡GÑ!ˆ
�A�t˜TØˆÜ47¸¿¹ÀdÇmÁmÐUY×UaÑUaÐcg×cpÑcp×4qÑ0ˆG�X˜w¨Ø—[‘[  A q¨!Ó,×7Ñ7Ó9×>Ñ>¸zÈDÑ?PÐRVÐXZÓ[×fÑfÓhˆFÙ˜f“oˆGØ—l‘l :¨t°T¸2Ó>×FÑFÀqÈ!ÈQÐPQÓR×]Ñ]Ó_ˆGÙ˜wÓ'ˆGØ˜‘-ˆCà—[‘[  A q¨!Ó,×7Ñ7Ó9×>Ñ>¸zÈDÑ?PÐRVÐXZÓ[×fÑfÓhˆFÙ˜f“oˆGØ—l‘l :¨t°T¸2Ó>×FÑFÀqÈ!ÈQÐPQÓR×]Ñ]Ó_ˆGÙ˜wÓ'ˆGØ˜‘-ŠCð 5rð �i‰i˜‹nˆØ�l‰l˜3 Ó%ˆØ�m‰m˜A˜qÓ!×,Ñ,Ó.ˆØˆ
r   )é@   é€   é   ÚLSTMé   éd   é2   )r%   r&   r'   r=   r>   rO   r   r)   r*   r   rk   ru   ry   r#   r+   r,   s   @r   rQ   rQ   `   sù   ø„ ñ
ð ØØØØØØñ)àð)ð ð)ð ð	)ð
 ð)ð ð)ð ð)ð ð)ð 
õ)ð<˜5Ÿ<™<ð ¨E°%·,±,ÀÐ2CÑ,Dó ð	˜%Ÿ,™,ð 	¨5°·±¸sÐ1BÑ+Có 	ð	˜Ÿ™ð 	¨Sð 	°U·\±\ó 	ð˜Ÿ™ð ¨%¯,©,÷ r   rQ   c                   ób   ‡ — e Zd Zddeddfˆ fd„Zdej                  dej                  fd„Zˆ xZS )ÚAutoPoolÚpool_dimr   Nc                 óÞ   •— t         t        | �  «        || _        t	        j
                  |¬«      | _        | j                  dt	        j                  t        j                  d«      «      «       y )Nr7   Úalphar	   )r   r‹   r   rŒ   r   ÚSoftmaxÚsoftmaxÚregister_parameterÚ	Parameterr)   Úones)r   rŒ   r   s     €r   r   zAutoPool.__init__Á   sH   ø€ ÜŒh˜Ñ&Ô(Ø%ˆŒÜ*,¯*©*¸Ô*BˆŒØ×Ñ ¬¯©´e·j±jÀ³mÓ)DÕEr   r   c                 óÎ   — | j                  t        j                  || j                  «      «      }t        j                  t        j                  ||«      | j
                  ¬«      }|S )Nr7   )r�   r)   ÚmulrŽ   ÚsumrŒ   )r   r   Úweightr"   s       r   r#   zAutoPool.forwardÇ   sC   € Ø—‘œeŸi™i¨¨4¯:©:Ó6Ó7ˆÜ�i‰iœŸ	™	 ! VÓ,°$·-±-Ô@ˆØˆ
r   )r	   )	r%   r&   r'   r>   r   r)   r*   r#   r+   r,   s   @r   r‹   r‹   À   s4   ø„ ñF ð F¨Tõ Fð˜Ÿ™ð ¨%¯,©,÷ r   r‹   c                   óª   ‡ — e Zd ZdZdej
                  dej
                  dej                  fˆ fd„Zdej                  de
ej                     fd„Zˆ xZS )	ÚSquimObjectiveaÙ  Speech Quality and Intelligibility Measures (SQUIM) model that predicts **objective** metric scores
    for speech enhancement (e.g., STOI, PESQ, and SI-SDR).

    Args:
        encoder (torch.nn.Module): Encoder module to transform 1D waveform to 2D feature representation.
        dprnn (torch.nn.Module): DPRNN module to model sequential feature.
        branches (torch.nn.ModuleList): Transformer branches in which each branch estimate one objective metirc score.
    ÚencoderÚdprnnÚbranchesc                 óT   •— t         t        | �  «        || _        || _        || _        y rM   )r   r™   r   rš   r›   rœ   )r   rš   r›   rœ   r   s       €r   r   zSquimObjective.__init__×   s'   ø€ ô 	Œn˜dÑ,Ô.ØˆŒØˆŒ
Ø ˆ�r   r   r   c                 óV  — |j                   dk7  rt        d|j                   › d�«      ‚|t        j                  |dz  dd¬«      dz  dz  z  }| j	                  |«      }| j                  |«      }g }| j                  D ])  }|j                   ||«      j                  d¬	«      «       Œ+ |S )
zá
        Args:
            x (torch.Tensor): Input waveforms. Tensor with dimensions `(batch, time)`.

        Returns:
            List(torch.Tensor): List of score Tenosrs. Each Tensor is with dimension `(batch,)`.
        r   z/The input must be a 2D Tensor. Found dimension Ú.r	   T)r8   Úkeepdimg      à?é   r7   )	ÚndimÚ
ValueErrorr)   Úmeanrš   r›   rœ   r_   Úsqueeze)r   r   r"   ÚscoresÚbranchs        r   r#   zSquimObjective.forwardâ   sž   € ð �6‰6�QŠ;ÜÐNÈqÏvÉvÈhÐVWÐXÓYÐYØ”—‘˜A˜q™D a°Ô6¸#Ñ=ÀÑBÑCˆØ�l‰l˜1‹oˆØ�j‰j˜‹oˆØˆØ—m”mˆFØ�M‰M™& ›+×-Ñ-°!Ð-Ó4Õ5ð $àˆr   )r%   r&   r'   r=   r   ÚModulerY   r   r)   r*   r   r#   r+   r,   s   @r   r™   r™   Í   sU   ø„ ñð	!à—‘ð	!ð �y‰yð	!ð —-‘-õ		!ð˜Ÿ™ð ¨$¨u¯|©|Ñ*<÷ r   r™   rT   ÚnheadÚmetricc                 ó¬  — t        j                  | || dz  dd¬«      }t        «       }|dk(  r[t        j                  t        j                  | | «      t        j
                  «       t        j                  | d«      t        «       «      }n·|dk(  rat        j                  t        j                  | | «      t        j
                  «       t        j                  | d«      t        t        ¬«      «      }nQt        j                  t        j                  | | «      t        j
                  «       t        j                  | d«      «      }t        j                  |||«      S )	al  Create branch module after DPRNN model for predicting metric score.

    Args:
        d_model (int): The number of expected features in the input.
        nhead (int): Number of heads in the multi-head attention model.
        metric (str): The metric name to predict.

    Returns:
        (nn.Module): Returned module to predict corresponding metric score.
    é   r$   T)rD   rF   Ústoir	   Úpesq)r   )r   ÚTransformerEncoderLayerr‹   ra   rJ   rc   r   r   )rT   r©   rª   Úlayer1Úlayer2Úlayer3s         r   Ú_create_branchr³   õ   sõ   € ô ×'Ñ'¨°¸À!¹ÈSÐ^bÔc€FÜ‹Z€FØ�ÒÜ—‘Ü�I‰I�g˜wÓ'Ü�H‰H‹JÜ�I‰I�g˜qÓ!Ü‹Nó	
‰ð 
�6Ò	Ü—‘Ü�I‰I�g˜wÓ'Ü�H‰H‹JÜ�I‰I�g˜qÓ!Ü¤9Ô-ó	
‰ô %'§M¡M´"·)±)¸GÀWÓ2MÌrÏxÉxËzÔ[]×[dÑ[dÐelÐnoÓ[pÓ$qˆÜ�=‰=˜ ¨Ó0Ð0r   r/   r0   rR   rS   rA   rU   rV   c	           	      óÖ   — |€|dz  }t        | |«      }	t        | ||||||«      }
t        j                  t	        ||d«      t	        ||d«      t	        ||d«      g«      }t        |	|
|«      S )aÃ  Build a custome :class:`torchaudio.models.squim.SquimObjective` model.

    Args:
        feat_dim (int, optional): The feature dimension after Encoder module.
        win_len (int): Kernel size in the Encoder module.
        d_model (int): The number of expected features in the input.
        nhead (int): Number of heads in the multi-head attention model.
        hidden_dim (int): Hidden dimension in the RNN layer of DPRNN.
        num_blocks (int): Number of DPRNN layers.
        rnn_type (str): Type of RNN in DPRNN. Valid options are ["RNN", "LSTM", "GRU"].
        chunk_size (int): Chunk size of input for DPRNN.
        chunk_stride (int or None, optional): Stride of chunk input for DPRNN.
    r   r­   r®   Úsisdr)r.   rQ   r   rY   r³   r™   )r/   r0   rT   r©   rR   rS   rA   rU   rV   rš   r›   rœ   s               r   Úsquim_objective_modelr¶     s~   € ð0 ÐØ! Q‘ˆÜ�h Ó(€GÜ�(˜J¨
°H¸gÀzÐS_Ó`€EÜ�}‰}ä˜7 E¨6Ó2Ü˜7 E¨6Ó2Ü˜7 E¨7Ó3ð	
ó€Hô ˜' 5¨(Ó3Ð3r   c            
      ó(   — t        dddddddd¬«      S )zSBuild :class:`torchaudio.models.squim.SquimObjective` model with default arguments.r‡   rƒ   r¬   r   r†   éG   )r/   r0   rT   r©   rR   rS   rA   rU   )r¶   © r   r   Úsquim_objective_baserº   ;  s'   € ä ØØØØØØØØô	ð 	r   rM   )r
   Útypingr   r   r   r)   Útorch.nnr   Útorch.nn.functionalÚ
functionalr:   r(   r   r   Ú__annotations__r¨   r   r.   r@   rQ   r‹   r™   r>   rO   Úmodulesr³   r¶   rº   r¹   r   r   Ú<module>rÁ      sa  ðÜ ß (Ñ (ã Ý ß Ð ðJ˜uð J¨ó Jð ñ ˜CÓ ð	"€	ˆ5�˜�Ñó ô	�2—9‘9ô 	ôˆb�i‰iô ô6�—	‘	ô ô4]ˆB�I‰Iô ]ô@
ˆr�y‰yô 
ô%�R—Y‘Yô %ðP1˜Cð 1¨ð 1°Sð 1¸R¿Z¹Z×=NÑ=Nó 1ðR #'ñ#4Øð#4àð#4ð ð#4ð ð	#4ð
 ð#4ð ð#4ð ð#4ð ð#4ð ˜3‘-ð#4ð ó#4ðL˜nô r   