Ë
    óÿæi¥  ã            
       óÒ   — d dl mZ d dlZd dlmZ d dl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de	f
d„Zde	fd„Zy)é    )ÚTupleNc                   ód   ‡ — e Zd ZdZdedefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )ÚAttPoolz±Attention-Pooling module that estimates the attention score.

    Args:
        input_dim (int): Input feature dimension.
        att_dim (int): Attention Tensor dimension.
    Ú	input_dimÚatt_dimc                 ó–   •— t         t        | �  «        t        j                  |d«      | _        t        j                  ||«      | _        y )Né   )Úsuperr   Ú__init__ÚnnÚLinearÚlinear1Úlinear2©Úselfr   r   Ú	__class__s      €úw/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/torchaudio/models/squim/subjective.pyr   zAttPool.__init__   s4   ø€ ÜŒg�tÑ%Ô'ä—y‘y ¨AÓ.ˆŒÜ—y‘y ¨GÓ4ˆ�ó    ÚxÚreturnc                 óú   — | j                  |«      }|j                  dd«      }t        j                  j	                  |d¬«      }t        j                  ||«      j                  d«      }| j                  |«      }|S )zïApply attention and pooling.

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

        Returns:
            (torch.Tensor): Attention score with dimensions `(batch, att_dim)`.
        é   r	   ©Údim)	r   Ú	transposer   Ú
functionalÚsoftmaxÚtorchÚmatmulÚsqueezer   )r   r   Úatts      r   ÚforwardzAttPool.forward   sg   € ð �l‰l˜1‹oˆØ�m‰m˜A˜qÓ!ˆÜ�m‰m×#Ñ# C¨QÐ#Ó/ˆÜ�L‰L˜˜aÓ ×(Ñ(¨Ó+ˆØ�L‰L˜‹OˆØˆr   ©
Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úintr   r   ÚTensorr"   Ú__classcell__©r   s   @r   r   r      s6   ø„ ñð5 #ð 5°õ 5ð˜Ÿ™ð ¨%¯,©,÷ r   r   c                   ód   ‡ — e Zd ZdZdedefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )Ú	PredictorzÏPrediction module that apply pooling and attention, then predict subjective metric scores.

    Args:
        input_dim (int): Input feature dimension.
        att_dim (int): Attention Tensor dimension.
    r   r   c                 óZ   •— t         t        | �  «        t        ||«      | _        || _        y ©N)r
   r-   r   r   Úatt_pool_layerr   r   s      €r   r   zPredictor.__init__0   s&   ø€ ÜŒi˜Ñ'Ô)Ü% i°Ó9ˆÔØˆ�r   r   r   c                 óî   — | j                  |«      }t        j                  j                  |d¬«      }t	        j
                  dd| j                  |j                  ¬«      }||z  j                  d¬«      }|S )a  Predict subjective evaluation metric score.

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

        Returns:
            (torch.Tensor): Subjective metric score. Tensor with dimensions `(batch,)`.
        r	   r   r   é   )ÚstepsÚdevice)	r0   r   r   r   r   Úlinspacer   r4   Úsum)r   r   ÚBs      r   r"   zPredictor.forward5   sb   € ð ×Ñ Ó"ˆÜ�M‰M×!Ñ! !¨Ð!Ó+ˆÜ�N‰N˜1˜a t§|¡|¸A¿H¹HÔEˆØ�‰U�K‰K˜AˆKÓˆØˆr   r#   r+   s   @r   r-   r-   (   s6   ø„ ñð #ð °õ ð
˜Ÿ™ð ¨%¯,©,÷ 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j                  de	ej                  ej                  f   fd	„Z
dej                  dej                  fd
„Zˆ xZS )ÚSquimSubjectiveaP  Speech Quality and Intelligibility Measures (SQUIM) model that predicts **subjective** metric scores
    for speech enhancement (e.g., Mean Opinion Score (MOS)). The model is adopted from *NORESQA-MOS*
    :cite:`manocha2022speech` which predicts MOS scores given the input speech and a non-matching reference.

    Args:
        ssl_model (torch.nn.Module): The self-supervised learning model for feature extraction.
        projector (torch.nn.Module): Projection layer that projects SSL feature to a lower dimension.
        predictor (torch.nn.Module): Predict the subjective scores.
    Ú	ssl_modelÚ	projectorÚ	predictorc                 óT   •— t         t        | �  «        || _        || _        || _        y r/   )r
   r9   r   r:   r;   r<   )r   r:   r;   r<   r   s       €r   r   zSquimSubjective.__init__P   s%   ø€ ÜŒo˜tÑ-Ô/Ø"ˆŒØ"ˆŒØ"ˆ�r   ÚwaveformÚ	referencer   c                 óØ   — |j                   d   }|j                   d   }||k  r6||z  dz   }t        j                  t        |«      D �cg c]  }|‘Œ c}d¬«      }||dd…d|…f   fS c c}w )aÙ  Cut or pad the reference Tensor to make it aligned with waveform Tensor.

        Args:
            waveform (torch.Tensor): Input waveform for evaluation. Tensor with dimensions `(batch, time)`.
            reference (torch.Tensor): Non-matching clean reference. Tensor with dimensions `(batch, time_ref)`.

        Returns:
            (torch.Tensor, torch.Tensor): The aligned waveform and reference Tensors
                with same dimensions `(batch, time)`.
        éÿÿÿÿr	   r   N)Úshaper   ÚcatÚrange)r   r>   r?   Ú
T_waveformÚT_referenceÚnum_paddingÚ_s          r   Ú_align_shapeszSquimSubjective._align_shapesV   s{   € ð —^‘^ BÑ'ˆ
Ø—o‘o bÑ)ˆØ˜Ò#Ø$¨Ñ3°aÑ7ˆKÜŸ	™	´e¸KÔ6HÓ"IÑ6H°¢9Ð6HÑ"IÈqÔQˆIØ˜¢1 k z k >Ñ2Ð2Ð2ùò #Js   Á	A'c                 óJ  — | j                  ||«      \  }}| j                  | j                  j                  |«      d   d   «      }| j                  | j                  j                  |«      d   d   «      }t	        j
                  ||fd¬«      }| j                  |«      }d|z
  S )a‰  Predict subjective evaluation metric score.

        Args:
            waveform (torch.Tensor): Input waveform for evaluation. Tensor with dimensions `(batch, time)`.
            reference (torch.Tensor): Non-matching clean reference. Tensor with dimensions `(batch, time_ref)`.

        Returns:
            (torch.Tensor): Subjective metric score. Tensor with dimensions `(batch,)`.
        r   rA   r   r   é   )rI   r;   r:   Úextract_featuresr   rC   r<   )r   r>   r?   ÚconcatÚ
score_diffs        r   r"   zSquimSubjective.forwardh   s—   € ð #×0Ñ0°¸9ÓEÑˆ�)Ø—>‘> $§.¡.×"AÑ"AÀ(Ó"KÈAÑ"NÈrÑ"RÓSˆØ—N‘N 4§>¡>×#BÑ#BÀ9Ó#MÈaÑ#PÐQSÑ#TÓUˆ	Ü—‘˜I xÐ0°aÔ8ˆØ—^‘^ FÓ+ˆ
Ø�:‰~Ðr   )r$   r%   r&   r'   r   ÚModuler   r   r)   r   rI   r"   r*   r+   s   @r   r9   r9   E   s„   ø„ ñð# "§)¡)ð #¸¿	¹	ð #ÈbÏiÉiõ #ð3 e§l¡lð 3¸u¿|¹|ð 3ÐPUÐV[×VbÑVbÐdi×dpÑdpÐVpÑPqó 3ð$ §¡ð ¸¿¹÷ r   r9   Ússl_typeÚfeat_dimÚproj_dimr   r   c                 ó¤   —  t        t        j                  | «      «       }t        j                  ||«      }t        |dz  |«      }t        |||«      S )a£  Build a custome :class:`torchaudio.prototype.models.SquimSubjective` model.

    Args:
        ssl_type (str): Type of self-supervised learning (SSL) models.
            Must be one of ["wav2vec2_base", "wav2vec2_large"].
        feat_dim (int): Feature dimension of the SSL feature representation.
        proj_dim (int): Output dimension of projection layer.
        att_dim (int): Dimension of attention scores.
    r   )ÚgetattrÚ
torchaudioÚmodelsr   r   r-   r9   )rP   rQ   rR   r   r:   r;   r<   s          r   Úsquim_subjective_modelrW   z   sJ   € ð 5”œ
×)Ñ)¨8Ó4Ó6€IÜ—	‘	˜( HÓ-€IÜ˜( Q™,¨Ó0€IÜ˜9 i°Ó;Ð;r   c                  ó    — t        dddd¬«      S )zXBuild :class:`torchaudio.prototype.models.SquimSubjective` model with default arguments.Úwav2vec2_basei   é    rK   )rP   rQ   rR   r   )rW   © r   r   Úsquim_subjective_baser\   �   s   € ä!Ø ØØØô	ð r   )Útypingr   r   Útorch.nnr   rU   rO   r   r-   r9   Ústrr(   rW   r\   r[   r   r   Ú<module>r`      s‚   ðÝ ã Ý Û ôˆb�i‰iô ô@�—	‘	ô ô:2�b—i‘iô 2ðj<Øð<àð<ð ð<ð ð	<ð
 ó<ð*˜ô r   