
     i                     v    d dl mZ d dl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  G d d	e      Zy)
    )	lru_cache)OptionalN)	rearrange)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 Zdej                  dej                  fdZ xZS )SimpleSegmentationModel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       }/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/models/segmentation/debug.pyr   z SimpleSegmentationModel.__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"SimpleSegmentationModel.num_frames=   s}    ( YY--99DD
		((44::))55<<{j000e+
:::r%   r/   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"   r/   r,   r-   s       r$   receptive_field_sizez,SimpleSegmentationModel.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"   r2   r,   r-   r.   s        r$   receptive_field_centerz.SimpleSegmentationModel.receptive_field_centerl   sw     YY--99DD
		((44::))55<<:%%:%
22r%   c                     t        | j                  t              rt        d      | j                  j                  r| j                  j
                  S t        | j                  j                        S )zDimension of outputz7SimpleSegmentationModel does not support multi-tasking.)
isinstancespecificationstuple
ValueErrorpowersetnum_powerset_classeslenclassesr"   s    r$   	dimensionz!SimpleSegmentationModel.dimension   sX     d))51VWW''&&;;;t**2233r%   c                 x    t        j                  d| j                        | _        | j	                         | _        y )N@   )r   Linearr?   
classifierdefault_activation
activationr>   s    r$   buildzSimpleSegmentationModel.build   s*     ))FDNN;113r%   	waveformsc                     | j                  |      }| j                  t        |d            \  }}| j                  | j	                  |            S )z

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

        Returns
        -------
        scores : (batch, time, classes)
        zb c f t -> b t (c f))r   r!   r   rE   rC   )r"   rG   r   outputhiddens        r$   forwardzSimpleSegmentationModel.forward   sD     yy#9T3I#JKtv677r%   )i>  r   N)r   )r   )__name__
__module____qualname__intr   r   r   r   r/   r1   r4   propertyr?   rF   torchTensorrK   __classcell__)r#   s   @r$   r
   r
   $   s     !#	

 
 tn	
0 ;c ;c ; ;85s 53 5$3C 3 3. 43 4 448 8%,, 8r%   r
   )	functoolsr   typingr   rQ   torch.nnr   einopsr   torchaudio.transformsr   pyannote.audio.core.modelr   pyannote.audio.core.taskr   r
    r%   r$   <module>r\      s-   0       & + )@8e @8r%   