
     i                     z    d dl mZ d dlmZ d dlZd dlmZ d dlmZm	Z	 d dl
mZ d dlmZ d dlmZ  G d d	e      Zy)
    )	lru_cache)OptionalN)	rearrangereduce)MFCC)Model)Taskc                        e Zd Z	 	 	 ddededee   f f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
edefd       Zdej                  dej                  fdZ xZS )SimpleEmbeddingModelsample_ratenum_channelstaskc                    t         |   |||       t        | j                  j                  dddd      | _        t        j                  | j
                  j                  | j                  j                  z  ddd	d	
      | _
        y )N)r   r   r   (      orthoF)r   n_mfccdct_typenormlog_mels       T)
num_layersbatch_firstbidirectional)super__init__r   hparamsr   mfccnnLSTMr   r   lstm)selfr   r   r   	__class__s       z/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/models/embedding/debug.pyr   zSimpleEmbeddingModel.__init__%   sy     	[|RVW00
	 GGIIt||888
	    num_samplesreturnc                 (   | j                   j                  j                  j                  }| j                   j                  j                  j                  }| j                   j                  j                  j
                  }|rd||z  z   S d||z
  |z  z   S )a|  Compute number of output frames for a given number of input samples

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

        Returns
        -------
        num_frames : int
            Number of output frames

        Source
        ------
        https://pytorch.org/docs/stable/generated/torch.stft.html#torch.stft

        r   r   MelSpectrogramspectrogram
hop_lengthn_fftcenter)r#   r'   r-   r.   r/   s        r%   
num_frameszSimpleEmbeddingModel.num_frames=   s}    ( YY--99DD
		((44::))55<<{j000e+
:::r&   r0   c                     | j                   j                  j                  j                  }| j                   j                  j                  j                  }||dz
  |z  z   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   r+   r,   r-   r.   )r#   r0   r-   r.   s       r%   receptive_field_sizez)SimpleEmbeddingModel.receptive_field_sizeZ   sN     YY--99DD
		((44::
Q*444r&   framec                 "   | j                   j                  j                  j                  }| j                   j                  j                  j                  }| j                   j                  j                  j
                  }|r||z  S ||z  |dz  z   S )zCompute center of receptive field

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

        Returns
        -------
        receptive_field_center : int
            Index of receptive field center.
        r   r*   )r#   r3   r-   r.   r/   s        r%   receptive_field_centerz+SimpleEmbeddingModel.receptive_field_centerl   sw     YY--99DD
		((44::))55<<:%%:%
22r&   c                      y)zDimension of output@    )r#   s    r%   	dimensionzSimpleEmbeddingModel.dimension   s     r&   	waveformsc                 z    | j                  |      }| j                  t        |d            \  }}t        |dd      S )z

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

        Returns
        -------
        embedding : (batch, dimension)
        zb c f t -> b t (c f)zb t f -> b fmean)r   r"   r   r   )r#   r:   r   outputhiddens        r%   forwardzSimpleEmbeddingModel.forward   s;     yy#9T3I#JKfnf55r&   )i>  r   N)r   )r   )__name__
__module____qualname__intr   r	   r   r   r0   r2   r5   propertyr9   torchTensorr?   __classcell__)r$   s   @r%   r   r   $   s     !#	

 
 tn	
0 ;c ;c ; ;85s 53 5$3C 3 3. 3  6 6%,, 6r&   r   )	functoolsr   typingr   rE   torch.nnr    einopsr   r   torchaudio.transformsr   pyannote.audio.core.modelr   pyannote.audio.core.taskr	   r   r8   r&   r%   <module>rO      s-   0      $ & + )s65 s6r&   