
     i:-                         d dl mZ d dlmZmZ d dlZd dlmZ d dlmc m	Z
 d dl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  G d	 d
e      Zy)    )	lru_cache)OptionalUnionN)pairwise)Model)Task)
merge_dict)conv1d_num_framesconv1d_receptive_field_centerconv1d_receptive_field_sizec                   "    e Zd ZdZdZddddddZddd	Z	 	 	 	 	 	 	 	 dd
eee	f   de
dedee   dee   dedede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 )!	SSeRiouSSa  Self-Supervised Representation for Speaker Segmentation

    wav2vec > LSTM > Feed forward > Classifier

    Parameters
    ----------
    sample_rate : int, optional
        Audio sample rate. Defaults to 16kHz (16000).
    num_channels : int, optional
        Number of channels. Defaults to mono (1).
    wav2vec: dict or str, optional
        Defaults to "WAVLM_BASE".
    wav2vec_frozen: bool, optional
        Whether to freeze wav2vec weights. Defaults to False.
    wav2vec_layer: int, optional
        Index of layer to use as input to the LSTM.
        Defaults (-1) to use average of all layers (with learnable weights).
    lstm : dict, optional
        Keyword arguments passed to the LSTM layer.
        Defaults to {"hidden_size": 128, "num_layers": 4, "bidirectional": True},
        i.e. two bidirectional layers with 128 units each.
        Set "monolithic" to False to split monolithic multi-layer LSTM into multiple mono-layer LSTMs.
        This may proove useful for probing LSTM internals.
    linear : dict, optional
        Keyword arugments used to initialize linear layers
        Defaults to {"hidden_size": 128, "num_layers": 2},
        i.e. two linear layers with 128 units each.
    
WAVLM_BASE      T        )hidden_size
num_layersbidirectional
monolithicdropout   )r   r   wav2vecwav2vec_frozenwav2vec_layerlstmlinearsample_ratenum_channelstaskc	           
      \   t         |   |||       t        |t              rt	        t
        j                  |      ryt        t
        j                  |      }	||	j                  k7  rt        d|	j                   d| d      |	j                  d   }
|	j                  d   }|	j                         | _        nt        j                  |      }|j                  d      }t        j                   j"                  di || _        |j                  d      }| j                  j%                  |       |d   }
|d   }n>t        |t&              r.t        j                   j"                  di || _        |d   }
|d   }|d	k  r/t)        j*                  t        j,                        d
      | _        | j                  j1                         D ]
  }| |_         t5        | j6                  |      }d
|d<   t5        | j8                  |      }| j;                  ddddd       |d   }|r*t'        |      }|d= t)        j<                  
fi || _        n|d   }|dkD  rt)        j@                  |d         | _!        t'        |      }d|d<   d|d<   |d= t)        jD                  tG        |      D cg c],  }t)        j<                  |d	k(  r
n|d   |d   rdndz  fi |. c}      | _        |d   dk  ry | jH                  j>                  d   | jH                  j>                  d   rdndz  }t)        jD                  tK        |g| jH                  jL                  d   g| jH                  jL                  d   z  z         D cg c]  \  }}t)        jN                  ||       c}}      | _&        y c c}w c c}}w )N)r   r   r    z	Expected z
Hz, found zHz.encoder_embed_dimencoder_num_layersconfig
state_dictr   T)datarequires_gradbatch_firstr   r   r   r   r   r   r      r   )pr   r   r   r    )(super__init__
isinstancestrhasattr
torchaudio	pipelinesgetattr_sample_rate
ValueError_params	get_modelr   torchloadpopmodelswav2vec2_modelload_state_dictdictnn	Parameteroneswav2vec_weights
parametersr'   r	   LSTM_DEFAULTSLINEAR_DEFAULTSsave_hyperparametersLSTMr   Dropoutr   
ModuleListrangehparamsr   r   Linear)selfr   r   r   r   r   r   r   r    bundlewav2vec_dimwav2vec_num_layers_checkpointr%   paramr   multi_layer_lstmr   one_layer_lstmilstm_out_featuresin_featuresout_features	__class__s                          /Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/models/segmentation/SSeRiouSS.pyr-   zSSeRiouSS.__init__S   s    	[|RVWgs#z++W5 !5!5w?&"5"55$#F$7$7#8
;-sS  %nn-@A%+^^4H%I"%//1 $jj1%//(3)00??J'J(__\:
,,Z8%&9:%,-A%B" &%,,;;FgFDL!"56K!()=!>1#%<<ZZ 234$D  \\,,.E&4"4E / $,,d3"]D00&9!!'&(	
 ,'
#Dz .@/?@DI l+JA~!zzDO<!$ZN+,N<((+N9%|, #:. / GG  !Av (!%m!4$($9qq"B	 ) /DI ,!#!%!2!2=!A""?3A"
 mm 2:) ||**=9:ll)),7882	2-K 		+|42	
)*	s   %1N#5 N(
returnc                     t        | j                  t              rt        d      | j                  j                  r| j                  j
                  S t        | j                  j                        S )zDimension of outputz)SSeRiouSS does not support multi-tasking.)r.   specificationstupler5   powersetnum_powerset_classeslenclasses)rM   s    rZ   	dimensionzSSeRiouSS.dimension   sX     d))51HII''&&;;;t**2233    c                 R   | j                   j                  d   dkD  r| j                   j                  d   }n7| j                   j                  d   | j                   j                  d   rdndz  }t        j                  || j
                        | _        | j                         | _        y )Nr   r   r   r   r   r)   )	rK   r   r   r?   rL   rc   
classifierdefault_activation
activation)rM   rW   s     rZ   buildzSSeRiouSS.build   s    <<|,q0,,--m<K,,++M:\\&&7QK ))K@113rd   num_samplesc           	          |}| j                   j                  j                  D ]T  }t        ||j                  |j
                  |j                  j                  d   |j                  j                  d         }V |S )zCompute number of output frames

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

        Returns
        -------
        num_frames : int
            Number of output frames.
        r   kernel_sizestridepaddingdilation)	r   feature_extractorconv_layersr
   rm   rn   convro   rp   )rM   rj   
num_frames
conv_layers       rZ   rt   zSSeRiouSS.num_frames   so     !
,,88DDJ*&22!(("//2#11!4J E rd   rt   c           	      
   |}t        | j                  j                  j                        D ]T  }t	        ||j
                  |j                  |j                  j                  d   |j                  j                  d         }V |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   )rt   rm   rn   ro   rp   )
reversedr   rq   rr   r   rm   rn   rs   ro   rp   )rM   rt   receptive_field_sizeru   s       rZ   rx   zSSeRiouSS.receptive_field_size   sv      *"4<<#A#A#M#MNJ#>/&22!(("//2#11!4$  O $#rd   framec           	      
   |}t        | j                  j                  j                        D ]T  }t	        ||j
                  |j                  |j                  j                  d   |j                  j                  d         }V |S )zCompute center of receptive field

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

        Returns
        -------
        receptive_field_center : int
            Index of receptive field center.
        r   rl   )
rw   r   rq   rr   r   rm   rn   rs   ro   rp   )rM   ry   receptive_field_centerru   s       rZ   r{   z SSeRiouSS.receptive_field_center	  sv     "'"4<<#A#A#M#MNJ%B&&22!(("//2#11!4&" O &%rd   	waveformsc                 "   | j                   j                  dk  rdn| j                   j                  }| j                  j                  |j	                  d      |      \  }}|:t        j                  |d      t        j                  | j                  d      z  }n|d   }| j                   j                  d   r| j                  |      \  }}nYt        | j                        D ]A  \  }} ||      \  }}|dz   | j                   j                  d   k  s1| j                  |      }C | j                   j                  d   dkD  r,| j                  D ]  }t        j                   ||            } | j                  | j!                  |            S )	zPass forward

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

        Returns
        -------
        scores : (batch, frame, classes)
        r   Nr)   )r   )dimr   r   )rK   r   r   extract_featuressqueezer8   stackFsoftmaxrB   r   	enumerater   r   
leaky_relurh   rf   )rM   r|   r   outputs_rU   r   r   s           rZ   forwardzSSeRiouSS.forward!  sZ    LL..2D8R8R 	 \\22a Z 3 

 kk'r2QYY$$!6 G bkG<<\*7+JGQ$TYY/4!']
q54<<,,\::"ll73G 0
 <<|,q0++,,vg7 & tw788rd   )NFr~   NNi>  r)   N)r)   )r   )__name__
__module____qualname____doc__WAV2VEC_DEFAULTSrD   rE   r   r>   r/   boolintr   r   r-   propertyrc   ri   r   rt   rx   r{   r8   Tensorr   __classcell__)rY   s   @rZ   r   r   *   sC   : $ M '*;O %)$#!% #j
tSy!j
 j
 	j

 tnj
 j
 j
 j
 tnj
X 43 4 4	4 c c  4$s $3 $2&C & &0'9 '9%,, '9rd   r   )	functoolsr   typingr   r   r8   torch.nnr?   torch.nn.functional
functionalr   r1   pyannote.core.utils.generatorsr   pyannote.audio.core.modelr   pyannote.audio.core.taskr   pyannote.audio.utils.paramsr	   $pyannote.audio.utils.receptive_fieldr
   r   r   r   r+   rd   rZ   <module>r      s@   .   "      3 + ) 2 ^9 ^9rd   