Ë
    ÿÍ:j™2  ã                   ó  — d dl mZ d dlmZ d dlZd dlmZ d dlmc 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mZ 	 d d
lmZ d dlmZ dZ	 d dlmZ dZ  G d„ de«      Z!y# e$ r dZY Œw xY w# e$ r dZ Y Œ"w xY w)é    )Ú	lru_cache)ÚOptionalN)Úmake_enc_dec)Úpairwise)ÚModel)ÚTask)Ú
merge_dict)Úconv1d_num_framesÚconv1d_receptive_field_centerÚconv1d_receptive_field_size)ÚDPRNN)Ú
pad_x_to_yTF)Ú	AutoModelc                   ó@  ‡ — e Zd ZdZdddddœZdddœZd	d
d
dddddœZddiZ	 	 	 	 	 	 	 	 	 	 	 d)dede	e   de	e   dede
de
de	e   de
dededefˆ fd„Zede
fd„«       Zd „ Zed!e
de
fd"„«       Zd*d#e
de
fd$„Zd+d%e
de
fd&„Zd'ej.                  dej.                  fd(„Zˆ xZS ),Ú	ToTaToNetu—  ToTaToNet joint speaker diarization and speech separation model

                        /--------------\
    Conv1D Encoder --------+--- DPRNN --X------- Conv1D Decoder
    WavLM -- upsampling --/                 \--- Avg pool -- Linear -- Classifier


    Parameters
    ----------
    sample_rate : int, optional
        Audio sample rate. Defaults to 16kHz (16000).
    num_channels : int, optional
        Number of channels. Defaults to mono (1).
    sincnet : dict, optional
        Keyword arguments passed to the SincNet block.
        Defaults to {"stride": 1}.
    linear : dict, optional
        Keyword arguments used to initialize linear layers
        See ToTaToNet.LINEAR_DEFAULTS for default values.
    diar : dict, optional
        Keyword arguments used to initialize the average pooling in the diarization branch.
        See ToTaToNet.DIAR_DEFAULTS for default values.
    encoder_decoder : dict, optional
        Keyword arguments used to initialize the encoder and decoder.
        See ToTaToNet.ENCODER_DECODER_DEFAULTS for default values.
    dprnn : dict, optional
        Keyword arguments used to initialize the DPRNN model.
        See ToTaToNet.DPRNN_DEFAULTS for default values.
    sample_rate : int, optional
        Audio sample rate. Defaults to 16000.
    num_channels : int, optional
        Number of channels. Defaults to 1.
    task : Task, optional
        Task to perform. Defaults to None.
    n_sources : int, optional
        Number of separated sources. Defaults to 3.
    use_wavlm : bool, optional
        Whether to use the WavLM large model for feature extraction. Defaults to True.
    wavlm_frozen : bool, optional
        Whether to freeze the WavLM model. Defaults to False.
    gradient_clip_val : float, optional
        Gradient clipping value. Required when fine-tuning the WavLM model and thus using two different optimizers.
        Defaults to 5.0.

    References
    ----------
    Joonas Kalda, ClÃ©ment PagÃ©s, Ricard Marxer, Tanel AlumÃ¤e, and HervÃ© Bredin.
    "PixIT: Joint Training of Speaker Diarization and Speech Separation
    from Real-world Multi-speaker Recordings"
    Odyssey 2024. https://arxiv.org/abs/2403.02288
    Úfreeé    é@   é   )Úfb_nameÚkernel_sizeÚ	n_filtersÚstrideé   )Úhidden_sizeÚ
num_layersé   é€   éd   ÚgLNÚreluÚLSTM)Ú	n_repeatsÚbn_chanÚhid_sizeÚ
chunk_sizeÚ	norm_typeÚmask_actÚrnn_typeÚframes_per_secondé}   Úencoder_decoderÚlinearÚdiarÚdprnnÚsample_rateÚnum_channelsÚtaskÚ	n_sourcesÚ	use_wavlmÚwavlm_frozenÚgradient_clip_valc           
      ó®  •— t         st        d«      ‚t        st        d«      ‚t        ‰| �  |||¬«       t        | j                  |«      }t        | j                  |«      }t        | j                  |«      }t        | j                  |«      }|	| _
        | j                  ddddd«       || _        |d	   d
k(  r|d   }n+|d	   dk(  rt        d|d   dz  dz   z  «      }nt        d«      ‚t        dd|i| j                   j"                  ¤Ž\  | _        | _        | j                  �rt)        j*                  d«      | _        | j,                  j/                  «       D ]
  }|
 |_        Œ d}| j,                  j2                  j4                  D ]C  }t7        |j8                  t:        j<                  «      sŒ(||j8                  j>                  d   z  }ŒE t        ||d   z  «      | _         tC        |d   | j,                  jD                  jF                  jH                  z   f|d   |dœ| j                   jJ                  ¤Ž| _&        n.tC        |d   f|d   |dœ| j                   jJ                  ¤Ž| _&        t        ||d   z  |d   z  «      | _'        t;        jP                  | jN                  | jN                  ¬«      | _)        |}|d   dkD  r€t;        jT                  tW        |g| j                   jX                  d   g| j                   jX                  d   z  z   «      D ��cg c]  \  }}t;        jZ                  ||«      ‘Œ c}}«      | _,        || _.        |
| _/        y c c}}w )Nzw'asteroid' must be installed to use ToTaToNet separation. `pip install pyannote-audio[separation]` should do the trick.z{'transformers' must be installed to use ToTaToNet separation. `pip install pyannote-audio[separation]` should do the trick.)r0   r1   r2   r,   r-   r/   r.   r5   r   r   r   Ústftr   é   zFilterbank type not recognized.r0   zmicrosoft/wavlm-larger   r   )Úout_chanÚn_srcr*   )r   r   r   © )0ÚASTEROID_IS_AVAILABLEÚImportErrorÚTRANSFORMERS_IS_AVAILABLEÚsuperÚ__init__r	   ÚLINEAR_DEFAULTSÚDPRNN_DEFAULTSÚENCODER_DECODER_DEFAULTSÚDIAR_DEFAULTSr4   Úsave_hyperparametersr3   ÚintÚ
ValueErrorr   Úhparamsr,   ÚencoderÚdecoderr   Úfrom_pretrainedÚwavlmÚ
parametersÚrequires_gradÚfeature_extractorÚconv_layersÚ
isinstanceÚconvÚnnÚConv1dr   Úwavlm_scalingr   Úfeature_projectionÚ
projectionÚout_featuresr/   ÚmaskerÚdiarization_scalingÚ	AvgPool1dÚaverage_poolÚ
ModuleListr   r-   ÚLinearr6   Úautomatic_optimization)Úselfr,   r-   r.   r/   r0   r1   r2   r3   r4   r5   r6   Ún_feats_outÚparamÚdownsampling_factorÚ
conv_layerÚlinear_input_featuresÚin_featuresrY   Ú	__class__s                      €ú/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pyannote/audio/models/separation/ToTaToNet.pyrA   zToTaToNet.__init__ƒ   sg  ø€ õ %ÜðPóð õ
 )ÜðPóð ô
 	‰Ñ [¸|ÐRVÐÔWä˜D×0Ñ0°&Ó9ˆÜ˜4×.Ñ.°Ó6ˆÜ$ T×%BÑ%BÀOÓTˆÜ˜$×,Ñ,¨dÓ3ˆØ"ˆŒØ×!Ñ!Ø˜x¨°&¸.ô	
ð #ˆŒà˜9Ñ%¨Ò/Ø)¨+Ñ6‰KØ˜YÑ'¨6Ò1Ü˜a ?°;Ñ#?À!Ñ#CÀaÑ#GÑHÓI‰KäÐ>Ó?Ð?Ü%1ñ &
Ø#ð&
Ø'+§|¡|×'CÑ'Cñ&
Ñ"ˆŒ�d”lð �>‹>Ü"×2Ñ2Ð3JÓKˆDŒJØŸ™×.Ñ.Ó0ò 7�Ø*6Ð&6�Õ#ð7à"#ÐØ"Ÿj™j×:Ñ:×FÑFò E�
Ü˜jŸo™o¬r¯y©yÕ9Ø'¨:¯?©?×+AÑ+AÀ!Ñ+DÑDÑ'ðEô "%Ð%8¸?È8Ñ;TÑ%TÓ!UˆDÔäØ Ñ,Ø—*‘*×/Ñ/×:Ñ:×GÑGñHðð )¨Ñ5Øñ	ð
 —,‘,×$Ñ$ñˆD�Kô  Ø Ñ,ðà(¨Ñ5Øñð —,‘,×$Ñ$ñ	ˆDŒKô $'Ø˜$Ð2Ñ3Ñ3°oÀhÑ6OÑOó$
ˆÔ ô ŸL™LØ×$Ñ$¨T×-EÑ-Eô
ˆÔð !,ÐØ�,Ñ !Ò#ÜŸ-™-ô 6>à1ðð  Ÿ<™<×.Ñ.¨}Ñ=Ð>ØŸ,™,×-Ñ-¨lÑ;ñ<ñ<ó6÷	á1˜ \ô —I‘I˜k¨<Õ8ó	óˆDŒKð "3ˆÔà&2ˆÕ#ùó	s   Ì M
Úreturnc                  ó   — y)zDimension of outputr9   r<   ©ra   s    ri   Ú	dimensionzToTaToNet.dimensionå   s   € ð ó    c                 óü   — | j                   j                  d   dkD  r&t        j                  d| j                  «      | _        n%t        j                  d| j                  «      | _        | j                  «       | _        y )Nr   r   r   r9   )rI   r-   rT   r_   rm   Ú
classifierÚdefault_activationÚ
activationrl   s    ri   ÚbuildzToTaToNet.buildê   sU   € Ø�<‰<×Ñ˜|Ñ,¨qÒ0Ü Ÿi™i¨¨D¯N©NÓ;ˆD�Oä Ÿi™i¨¨4¯>©>Ó:ˆDŒOØ×1Ñ1Ó3ˆ�rn   Únum_samplesc                 ó¶   — | j                   | j                  j                  d   z  }| j                   | j                  j                  d   z  }t        |||¬«      S )zíCompute number of output frames

        Parameters
        ----------
        num_samples : int
            Number of input samples.

        Returns
        -------
        num_frames : int
            Number of output frames.
        r   r   ©r   r   )r[   rI   r,   r
   )ra   rt   Úequivalent_strideÚequivalent_kernel_sizes       ri   Ú
num_frameszToTaToNet.num_framesñ   sb   € ð  ×$Ñ$ t§|¡|×'CÑ'CÀHÑ'MÑMð 	ð ×$Ñ$ t§|¡|×'CÑ'CÀMÑ'RÑRð 	ô !ØÐ%;ÐDUô
ð 	
rn   ry   c                 ó¶   — | j                   | j                  j                  d   z  }| j                   | j                  j                  d   z  }t        |||¬«      S )a
  Compute size of receptive field

        Parameters
        ----------
        num_frames : int, optional
            Number of frames in the output signal

        Returns
        -------
        receptive_field_size : int
            Receptive field size.
        r   r   rv   )r[   rI   r,   r   )ra   ry   rw   rx   s       ri   Úreceptive_field_sizezToTaToNet.receptive_field_size  sb   € ð ×$Ñ$ t§|¡|×'CÑ'CÀHÑ'MÑMð 	ð ×$Ñ$ t§|¡|×'CÑ'CÀMÑ'RÑRð 	ô +ØÐ$:ÐCTô
ð 	
rn   Úframec                 ó¶   — | j                   | j                  j                  d   z  }| j                   | j                  j                  d   z  }t        |||¬«      S )zúCompute center of receptive field

        Parameters
        ----------
        frame : int, optional
            Frame index

        Returns
        -------
        receptive_field_center : int
            Index of receptive field center.
        r   r   rv   )r[   rI   r,   r   )ra   r|   rw   rx   s       ri   Úreceptive_field_centerz ToTaToNet.receptive_field_center$  sb   € ð ×$Ñ$ t§|¡|×'CÑ'CÀHÑ'MÑMð 	ð ×$Ñ$ t§|¡|×'CÑ'CÀMÑ'RÑRð 	ô -ØÐ5Ð>Oô
ð 	
rn   Ú	waveformsc                 óV  — |j                   d   }| j                  |«      }| j                  r�| j                  |j	                  d«      «      j
                  }|j                  dd«      }|j                  | j                  d¬«      }t        ||«      }t        j                  ||fd¬«      }| j                  |«      }n| j                  |«      }||j                  d«      z  }| j                  |«      }t        ||«      }|j                  dd«      }t        j                  |dd¬«      }| j!                  |«      }|j                  dd«      }| j"                  j$                  d   dkD  r,| j$                  D ]  }	t'        j(                   |	|«      «      }Œ | j"                  j$                  d   dk(  r$|dz  j+                  d¬«      j                  d«      }| j-                  |«      }|j/                  || j0                  d«      }|j                  dd«      } | j2                  d   |«      |fS )zàPass forward

        Parameters
        ----------
        waveforms : (batch, channel, sample)

        Returns
        -------
        scores : (batch, frame, classes)
        sources : (batch, sample, n_sources)
        r   r9   r   éÿÿÿÿ)Údim)Ú	start_dimÚend_dimr   )ÚshaperJ   r4   rM   ÚsqueezeÚlast_hidden_stateÚ	transposeÚrepeat_interleaverV   r   ÚtorchÚcatrZ   Ú	unsqueezerK   Úflattenr]   rI   r-   ÚFÚ
leaky_reluÚsumrp   Úreshaper3   rr   )
ra   r   ÚbszÚtf_repÚ	wavlm_repÚmasksÚmasked_tf_repÚdecoded_sourcesÚoutputsr-   s
             ri   ÚforwardzToTaToNet.forward=  së  € ð �o‰o˜aÑ ˆØ—‘˜iÓ(ˆØ�>Š>ØŸ
™
 9×#4Ñ#4°QÓ#7Ó8×JÑJˆIØ!×+Ñ+¨A¨qÓ1ˆIØ!×3Ñ3°D×4FÑ4FÈBÐ3ÓOˆIÜ" 9¨fÓ5ˆIÜŸ	™	 6¨9Ð"5¸1Ô=ˆIØ—K‘K 	Ó*‰Eà—K‘K Ó'ˆEà × 0Ñ 0°Ó 3Ñ3ˆØŸ,™, }Ó5ˆÜ$ _°iÓ@ˆØ)×3Ñ3°A°qÓ9ˆÜ—-‘- ¸ÀAÔFˆà×#Ñ# GÓ,ˆØ×#Ñ# A qÓ)ˆà�<‰<×Ñ˜|Ñ,¨qÒ0ØŸ+™+ò 8�ÜŸ,™,¡v¨g£Ó7‘ð8à�<‰<×Ñ˜|Ñ,°Ò1Ø ‘z×&Ñ&¨1Ð&Ó-×7Ñ7¸Ó;ˆGØ—/‘/ 'Ó*ˆØ—/‘/ # t§~¡~°rÓ:ˆØ×#Ñ# A qÓ)ˆà!ˆt�‰˜qÑ! 'Ó*¨OÐ;Ð;rn   )NNNNi€>  r9   Né   TFg      @)r9   )r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__rD   rB   rC   rE   Údictr   rG   r   ÚboolÚfloatrA   Úpropertyrm   rs   r   ry   r{   r~   rŠ   ÚTensorr™   Ú__classcell__)rh   s   @ri   r   r   <   sƒ  ø„ ñ2ðj ØØØñ	 Ðð ')¸Ñ:€OàØØØØØØñ€Nð )¨#Ð.€Mð !%Ø!%Ø#ØØ ØØ#ØØØ"Ø#&ñ`3àð`3ð ˜‘ð`3ð �t‰nð	`3ð
 ð`3ð ð`3ð ð`3ð �t‰nð`3ð ð`3ð ð`3ð ð`3ð !õ`3ðD ð˜3ò ó ðò4ð ð
 cð 
¨cò 
ó ð
ñ2
¨sð 
¸3ó 
ñ2
¨Cð 
¸ó 
ð2*< §¡ð *<°%·,±,÷ *<rn   r   )"Ú	functoolsr   Útypingr   rŠ   Útorch.nnrT   Útorch.nn.functionalÚ
functionalrŽ   Úasteroid_filterbanksr   Úpyannote.core.utils.generatorsr   Úpyannote.audio.core.modelr   Úpyannote.audio.core.taskr   Úpyannote.audio.utils.paramsr	   Ú$pyannote.audio.utils.receptive_fieldr
   r   r   Úasteroid.masknnr   Úasteroid.utils.torch_utilsr   r=   r>   Útransformersr   r?   r   r<   rn   ri   ú<module>r³      sŒ   ðõ2  Ý ã Ý ß Ð Ý -Ý 3å +Ý )Ý 2÷ñ ð"Ý%Ý5à Ðð
&Ý&à $Ðô
k<�õ k<øð ò "Ø!Òð"ûð ò &Ø %Òð&ús$   Á	A, ÁA9 Á,A6Á5A6Á9BÂB