
     iNr                     j   d dl Z d dlmZ d dlmZ d dlmZmZmZ d dl	Z
d dlZd dlmc 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mZmZ d d	lmZ d d
lm Z  d dl!m"Z"m#Z# 	 d dl$m%Z& dZ'	 d dl)m*Z+ dZ,	 d dl-Z.dZ/ G d de      Z0 G d de      Z1 G d de      Z2 G d de      Z3	 	 	 d$de"deejh                     deedf   deeedf   fdZ5 G d de      Z6	 	 	 	 d%de7d e7de7d!ee7   fd"Z8e9d#k(  rd dl:Z: e:jv                  e8       yy# e($ r dZ'Y w xY w# e($ r dZ,Y w xY w# e($ r dZ/Y w xY w)&    N)cached_property)Path)OptionalTextUnion)hf_hub_download)RepositoryNotFoundError)pad_sequence)	InferenceModelPipeline)BaseInference)	AudioFile)PipelineModel	get_model)EncoderClassifierTF)EncDecSpeakerLabelModelc                       e Zd Z	 	 ddedeej                     f fdZdej                  fdZe	de
fd       Ze	de
fd       Ze	defd       Ze	de
fd	       Z	 dd
ej                   deej                      dej$                  fdZ xZS )NeMoPretrainedSpeakerEmbedding	embeddingdevicec                 \   t         st        d| d      t        |           || _        |xs t        j                  d      | _        t        j                  | j                        | _	        | j                  j                          | j                  j                  | j                         y )Nz!'NeMo' must be installed to use 'zQ' embeddings. Visit https://nvidia.github.io/NeMo/ for installation instructions.cpu)NEMO_IS_AVAILABLEImportErrorsuper__init__r   torchr   NeMo_EncDecSpeakerLabelModelfrom_pretrainedmodel_freezeto)selfr   r   	__class__s      /Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/pipelines/speaker_verification.pyr   z'NeMoPretrainedSpeakerEmbedding.__init__B   s    
 !3I; ?V V 
 	"3U 32BB4>>Rt{{#    c                     t        |t        j                        s"t        dt	        |      j
                   d      | j                  j                  |       || _        | S N5`device` must be an instance of `torch.device`, got ``
isinstancer   r   	TypeErrortype__name__r!   r#   r$   r   s     r&   r#   z!NeMoPretrainedSpeakerEmbedding.toU   S    &%,,/GVH]H]G^^_`  	vr'   returnc                 b    | j                   j                  j                  j                  dd      S )Nsample_rate>  )r!   _cfgtrain_dsgetr$   s    r&   r5   z*NeMoPretrainedSpeakerEmbedding.sample_rate_   s%    {{((,,]EBBr'   c                 4   t        j                  d| j                        j                  | j                        }t        j
                  | j                  g      j                  | j                        }| j                  ||      \  }}|j                  \  }}|S )N   input_signalinput_signal_length)r   randr5   r#   r   tensorr!   shape)r$   r>   r?   _
embeddings	dimensions         r&   rE   z(NeMoPretrainedSpeakerEmbedding.dimensionc   s    zz!T%5%5699$++F#llD,<,<+=>AA$++N%;N $ 
: "''9r'   c                      yNcosine r:   s    r&   metricz%NeMoPretrainedSpeakerEmbedding.metricm       r'   c                    dt        d| j                  z        }}||z   dz  }|dz   |k  r	 t        j                  d|      j	                  | j
                        }t        j                  |g      j	                  | j
                        }| j                  ||      }|}||z   dz  }|dz   |k  r|S # t        $ r |}Y  w xY w)N         ?r<   r=   )	roundr5   r   r@   r#   r   rA   r!   RuntimeError)r$   loweruppermiddler>   r?   rC   s          r&   min_num_samplesz.NeMoPretrainedSpeakerEmbedding.min_num_samplesq   s    %d&6&6 67u%-A%ai%
$zz!V477D&+llF8&<&?&?&L#KK!-CV     em)F ai%    s   A3B2 2C ?C 	waveformsmasksc                    |j                   \  }}}|dk(  sJ |j                  d      }|8|j                  d      }|j                   d   t        j                  |      z  }n|j                   \  }}	||k(  sJ t	        j
                  |j                  d      |d      j                  d      }
|
dkD  }
t        t        ||
      D cg c]
  \  }}||    c}}d      }|
j                  d      }|j                         }|| j                  k  r2t        j                  t        j                  || j                  f      z  S || j                  k  }|||<   | j!                  |j#                  | j$                        |j#                  | j$                              \  }	}|j'                         j)                         }t        j                  ||j'                         j)                         <   |S c c}}w )	   

        Parameters
        ----------
        waveforms : (batch_size, num_channels, num_samples)
            Only num_channels == 1 is supported.
        masks : (batch_size, num_samples), optional

        Returns
        -------
        embeddings : (batch_size, dimension)

        r<   dimnearestsizemoderN   Tbatch_firstr=   )rB   squeezer   onesFinterpolate	unsqueezer
   zipsummaxrT   npnanzerosrE   r!   r#   r   r   numpyr$   rU   rV   
batch_sizenum_channelsnum_samplessignalswav_lensbatch_size_masksrC   imaskswaveformimaskmax_len	too_shortrD   s                   r&   __call__z'NeMoPretrainedSpeakerEmbedding.__call__   s   " 1:-
L+q   %%!%,	=''A'.G}}Q'%**Z*@@H #(++a!1111 ]]A&[yg!gn  c\F"8;Iv8NO8N_Xu%8NO G
 zzaz(H,,. T)))66BHHj$..%ABBBt333	%"dkk2 (DKK 8 $ 
:
  ^^%++-
.0ff
9==?((*+/ Ps   G2
)z+nvidia/speakerverification_en_titanet_largeNN)r0   
__module____qualname__r   r   r   r   r   r#   r   intr5   rE   strrJ   rT   Tensorri   ndarrayry   __classcell__r%   s   @r&   r   r   A   s     H)-$$ &$&  CS C C 3        * HLAA.6u||.DA	Ar'   r   c                   @    e Zd ZdZ	 	 	 	 ddedeej                     deedf   dee	edf   f fdZ
dej                  fdZed	efd
       Zed	efd       Zed	efd       Zed	efd       Z	 ddej&                  deej&                     d	ej*                  fdZ xZS )%SpeechBrainPretrainedSpeakerEmbeddinga  Pretrained SpeechBrain speaker embedding

    Parameters
    ----------
    embedding : str
        Name of SpeechBrain model
    device : torch.device, optional
        Device
    token : str or bool, optional
        Huggingface token to be used for downloading from Huggingface hub.
    cache_dir: Path or str, optional
        Path to the folder where files downloaded from Huggingface hub are stored.

    Usage
    -----
    >>> get_embedding = SpeechBrainPretrainedSpeakerEmbedding("speechbrain/spkrec-ecapa-voxceleb")
    >>> assert waveforms.ndim == 3
    >>> batch_size, num_channels, num_samples = waveforms.shape
    >>> assert num_channels == 1
    >>> embeddings = get_embedding(waveforms)
    >>> assert embeddings.ndim == 2
    >>> assert embeddings.shape[0] == batch_size

    >>> assert binary_masks.ndim == 1
    >>> assert binary_masks.shape[0] == batch_size
    >>> embeddings = get_embedding(waveforms, masks=binary_masks)
    Nr   r   token	cache_dirc                    t         st        d| d      t        |           d|v r3|j	                  d      d   | _        |j	                  d      d   | _        n|| _        d | _        |xs t        j                  d      | _        || _	        || _
        t        j                  | j
                  | j                   dd| j                  i| j                  | j                  | j                  	      | _        y )
Nz('speechbrain' must be installed to use 'zP' embeddings. Visit https://speechbrain.github.io for installation instructions.@r   r<   r   /speechbrainr   sourcesavedirrun_optsr   huggingface_cache_dirrevision)SPEECHBRAIN_IS_AVAILABLEr   r   r   splitr   r   r   r   r   r   SpeechBrain_EncoderClassifierfrom_hparamsclassifier_r$   r   r   r   r   r%   s        r&   r   z.SpeechBrainPretrainedSpeakerEmbedding.__init__   s     (:9+ FU U 
 	)&__S1!4DN%OOC03DM&DN DM3U 3
"8EE>>~~&l3,**"&..]]
r'   c                 :   t        |t        j                        s"t        dt	        |      j
                   d      t        j                  | j                  | j                   dd|i| j                  | j                  | j                        | _        || _        | S )Nr*   r+   r   r   r   )r-   r   r   r.   r/   r0   r   r   r   r   r   r   r   r1   s     r&   r#   z(SpeechBrainPretrainedSpeakerEmbedding.to  s    &%,,/GVH]H]G^^_`  9EE>>~~&l3'**"&..]]
 r'   r3   c                 B    | j                   j                  j                  S rz   )r   audio_normalizerr5   r:   s    r&   r5   z1SpeechBrainPretrainedSpeakerEmbedding.sample_rate  s    00<<<r'   c                     t        j                  dd      j                  | j                        }| j                  j                  |      j                  ^ }}|S )Nr<   r6   )r   r@   r#   r   r   encode_batchrB   )r$   dummy_waveformsrC   rE   s       r&   rE   z/SpeechBrainPretrainedSpeakerEmbedding.dimension  sG    **Q.11$++>((55oFLLIr'   c                      yrG   rI   r:   s    r&   rJ   z,SpeechBrainPretrainedSpeakerEmbedding.metric#  rK   r'   c                    t        j                         5  dt        d| j                  z        }}||z   dz  }|dz   |k  r\	 | j                  j                  t        j                  d|      j                  | j                              }|}||z   dz  }|dz   |k  r\d d d        |S # t        $ r |}Y (w xY w# 1 sw Y   S xY wNrM   rN   r<   )
r   inference_moderO   r5   r   r   randnr#   r   rP   r$   rQ   rR   rS   rC   s        r&   rT   z5SpeechBrainPretrainedSpeakerEmbedding.min_num_samples'  s    !!#eC$*:*:$:;5Eem)F!)e##((55Av.11$++>A #E  %-A- !)e# $  $ #"E# $ s0   +B7A
B&B7&B41B73B44B77CrU   rV   c                    |j                   \  }}}|dk(  sJ |j                  d      }|8|j                  d      }|j                   d   t        j                  |      z  }n|j                   \  }}	||k(  sJ t	        j
                  |j                  d      |d      j                  d      }
|
dkD  }
t        t        ||
      D cg c]  \  }}||   j                          c}}d      }|
j                  d      }|j                         }|| j                  k  r2t        j                  t        j                  || j                   f      z  S || j                  k  }||z  }d||<   | j"                  j%                  ||	      j                  d      j'                         j)                         }t        j                  ||j'                         j)                         <   |S c c}}w )
rX   r<   rY   r[   r\   rN   Tr_   g      ?)rr   )rB   ra   r   rb   rc   rd   re   r
   rf   
contiguousrg   rh   rT   ri   rj   rk   rE   r   r   r   rl   rm   s                   r&   ry   z.SpeechBrainPretrainedSpeakerEmbedding.__call__9  s   " 1:-
L+q   %%!%,	=''A'.G}}Q'%**Z*@@H #(++a!1111 ]]A&[yg!gn  c\F" ,/y&+A+A% UO..0+A !G zzaz(H,,. T)))66BHHj$..%ABBBt333	g%! ))'H)EWW^SUUW	 	 /1ff
9==?((*+9s   G(
)z!speechbrain/spkrec-ecapa-voxcelebNNNrz   )r0   r{   r|   __doc__r   r   r   r   r   r   r   r#   r   r}   r5   rE   r~   rJ   rT   r   ri   r   ry   r   r   s   @r&   r   r      s   < >)-#'-1

 &
 T4Z 	

 tT)*
B " =S = = 3  
      $ HLFF.6u||.DF	Fr'   r   c                       e Zd ZdZ	 	 	 	 ddedeej                     deedf   dee	edf   f fdZ
dej                  fdZed	efd
       Zed	efd       Zed	efd       Zed	efd       Zed	efd       Z	 	 	 	 ddej(                  dedededed	ej(                  fdZ	 ddej(                  deej(                     d	ej0                  fdZ xZS )'ONNXWeSpeakerPretrainedSpeakerEmbeddinga  Pretrained WeSpeaker speaker embedding

    Parameters
    ----------
    embedding : str
        Path to WeSpeaker pretrained speaker embedding
    device : torch.device, optional
        Device
    token : str or bool, optional
        Huggingface token to be used for downloading from Huggingface hub.
    cache_dir: Path or str, optional
        Path to the folder where files downloaded from Huggingface hub are stored.

    Usage
    -----
    >>> get_embedding = ONNXWeSpeakerPretrainedSpeakerEmbedding("hbredin/wespeaker-voxceleb-resnet34-LM")
    >>> assert waveforms.ndim == 3
    >>> batch_size, num_channels, num_samples = waveforms.shape
    >>> assert num_channels == 1
    >>> embeddings = get_embedding(waveforms)
    >>> assert embeddings.ndim == 2
    >>> assert embeddings.shape[0] == batch_size

    >>> assert binary_masks.ndim == 1
    >>> assert binary_masks.shape[0] == batch_size
    >>> embeddings = get_embedding(waveforms, masks=binary_masks)
    Nr   r   r   r   c                 4   t         st        d| d      t        |           t	        |      j                         s	 t        |d||      }|| _	        | j                  |xs t        j                  d             y # t        $ r t        d| d      w xY w)Nz('onnxruntime' must be installed to use 'z' embeddings.zspeaker-embedding.onnx)repo_idfilenamer   r   zCould not find 'z&' on huggingface.co nor on local disk.r   )ONNX_IS_AVAILABLEr   r   r   r   existsr   r	   
ValueErrorr   r#   r   r   r   s        r&   r   z0ONNXWeSpeakerPretrainedSpeakerEmbedding.__init__  s     !:9+]S  	I%%'
+%5'		 #-%,,u-. +  &yk1WX s   A> >Bc                    t        |t        j                        s"t        dt	        |      j
                   d      |j                  dk(  rdg}nR|j                  dk(  rdddifg}n;t        j                  d	|j                   d
       t        j                  d      }dg}t        j                         }d|_
        d|_        t        j                  | j                  ||      | _        || _        | S )Nr*   r+   r   CPUExecutionProvidercudaCUDAExecutionProvidercudnn_conv_algo_searchDEFAULTzUnsupported device type: z, falling back to CPUr<   )sess_options	providers)r-   r   r   r.   r/   r0   warningswarnortSessionOptionsinter_op_num_threadsintra_op_num_threadsInferenceSessionr   session_)r$   r   r   r   s       r&   r#   z*ONNXWeSpeakerPretrainedSpeakerEmbedding.to  s    &%,,/GVH]H]G^^_`  ;;%/0I[[F" ,0)I MM+FKK=8MN \\%(F/0I))+,-),-),,NN
 r'   r3   c                      y)Nr6   rI   r:   s    r&   r5   z3ONNXWeSpeakerPretrainedSpeakerEmbedding.sample_rate  s    r'   c                     t        j                  ddd      }| j                  |      }| j                  j	                  dgd|j                         i      d   }|j                  \  }}|S )Nr<   r6   embsfeatsoutput_names
input_feedr   )r   r@   compute_fbankr   runrl   rB   )r$   r   featuresrD   rC   rE   s         r&   rE   z1ONNXWeSpeakerPretrainedSpeakerEmbedding.dimension  so    **Q51%%o6]]&& w8H.I ' 


 "''9r'   c                      yrG   rI   r:   s    r&   rJ   z.ONNXWeSpeakerPretrainedSpeakerEmbedding.metric  rK   r'   c                    dt        d| j                  z        }}||z   dz  }|dz   |k  r	 | j                  t        j                  dd|            }| j                  j                  dgd|j                         i      d   }t        j                  t        j                  |            r|}n|}||z   dz  }|dz   |k  r|S # t
        $ r |}||z   dz  }Y w xY w)NrM   rN   r<   r   r   r   r   )rO   r5   r   r   r   AssertionErrorr   r   rl   ri   anyisnan)r$   rQ   rR   rS   r   rD   s         r&   rT   z7ONNXWeSpeakerPretrainedSpeakerEmbedding.min_num_samples  s    %d&6&6 67u%-A%ai%--ekk!Q.GH **$X7HNN<L2M + J vvbhhz*+em)F# ai%&  " %-A-s   &C CCc                 |    | j                  t        j                  dd| j                              j                  d   S )Nr<   )r   r   r   rT   rB   r:   s    r&   min_num_framesz6ONNXWeSpeakerPretrainedSpeakerEmbedding.min_num_frames  s2    !!%++aD4H4H"IJPPQRSSr'   rU   num_mel_binsframe_lengthframe_shiftditherc                     |dz  }t        j                  |D cg c])  }t        j                  |||||| j                  dd      + c}      }|t        j
                  |dd      z
  S c c}w )af  Extract fbank features

        Parameters
        ----------
        waveforms : (batch_size, num_channels, num_samples)

        Returns
        -------
        fbank : (batch_size, num_frames, num_mel_bins)

        Source: https://github.com/wenet-e2e/wespeaker/blob/45941e7cba2c3ea99e232d02bedf617fc71b0dad/wespeaker/bin/infer_onnx.py#L30C1-L50
        i   hammingF)r   r   r   r   sample_frequencywindow_type
use_energyr<   T)rZ   keepdim)r   stackkaldifbankr5   mean)r$   rU   r   r   r   r   ru   r   s           r&   r   z5ONNXWeSpeakerPretrainedSpeakerEmbedding.compute_fbank  s    * )	;; !* !*H !-!- +!%)%5%5 )$	 !*
  %**X1dCCCs   .A)rV   c                    |j                   \  }}}|dk(  sJ | j                  |j                  | j                              }|j                   \  }}}|5| j                  j                  dgd|j                  d      i      d   }	|	S |j                   \  }
}||
k(  sJ t        j                  |j                  d	      |d
      j                  d	      }|dkD  }t        j                  t        j                  || j                  f      z  }	t        t!        ||            D ]f  \  }\  }}||   }|j                   d   | j"                  k  r+| j                  j                  dgd|j                  d      d   i      d   d   |	|<   h |	S )rX   r<   Nr   r   T)forcer   r   rY   r[   r\   rN   )rB   r   r#   r   r   r   rl   rc   rd   re   ra   ri   rj   rk   rE   	enumeraterf   r   )r$   rU   rV   rn   ro   rp   r   rC   
num_framesrD   rs   rt   ffeaturerv   masked_features                   r&   ry   z0ONNXWeSpeakerPretrainedSpeakerEmbedding.__call__8  s   " 1:-
L+q   %%ill4;;&?@#>>:q=**$X7HNNQUN<V2W + J #kk!----OOO")

'a'. 	 #VVbhh
DNN'CDD
#,S6-B#CA$U^N##A&)<)<< MM--$X#^%9%9%9%Ed%KL .   JqM $D r'   )z&hbredin/wespeaker-voxceleb-resnet34-LMNNN)P      
           rz   )r0   r{   r|   r   r   r   r   r   r   r   r   r#   r   r}   r5   rE   r~   rJ   rT   r   r   floatr   ri   r   ry   r   r   s   @r&   r   r     s   < C)-#'-1// &/ T4Z 	/
 tT)*/@   D S   3        0 T T T &D<<&D &D 	&D
 &D &D 
&DR HL33.6u||.D3	3r'   r   c                   @    e Zd ZdZ	 	 	 	 ddedeej                     dee	df   dee
e	df   f fdZdej                  fdZed	efd
       Zed	efd       Zed	efd       Zed	efd       Z	 ddej(                  deej(                     d	ej,                  fdZ xZS )'PyannoteAudioPretrainedSpeakerEmbeddinga  Pretrained pyannote.audio speaker embedding

    Parameters
    ----------
    embedding : PipelineModel
        pyannote.audio model
    device : torch.device, optional
        Device
    token : str or bool, optional
        Huggingface token to be used for downloading from Huggingface hub.
    cache_dir: Path or str, optional
        Path to the folder where files downloaded from Huggingface hub are stored.

    Usage
    -----
    >>> get_embedding = PyannoteAudioPretrainedSpeakerEmbedding("pyannote/embedding")
    >>> assert waveforms.ndim == 3
    >>> batch_size, num_channels, num_samples = waveforms.shape
    >>> assert num_channels == 1
    >>> embeddings = get_embedding(waveforms)
    >>> assert embeddings.ndim == 2
    >>> assert embeddings.shape[0] == batch_size

    >>> assert masks.ndim == 1
    >>> assert masks.shape[0] == batch_size
    >>> embeddings = get_embedding(waveforms, masks=masks)
    Nr   r   r   r   c                 $   t         |           || _        |xs t        j                  d      | _        t        | j                  ||      | _        | j                  j                          | j                  j                  | j                         y )Nr   r   r   )	r   r   r   r   r   r   r!   evalr#   r   s        r&   r   z0PyannoteAudioPretrainedSpeakerEmbedding.__init__  sd     	"3U 3&t~~UiXt{{#r'   c                     t        |t        j                        s"t        dt	        |      j
                   d      | j                  j                  |       || _        | S r)   r,   r1   s     r&   r#   z*PyannoteAudioPretrainedSpeakerEmbedding.to  r2   r'   r3   c                 B    | j                   j                  j                  S rz   )r!   audior5   r:   s    r&   r5   z3PyannoteAudioPretrainedSpeakerEmbedding.sample_rate  s    {{  ,,,r'   c                 .    | j                   j                  S rz   )r!   rE   r:   s    r&   rE   z1PyannoteAudioPretrainedSpeakerEmbedding.dimension  s    {{$$$r'   c                      yrG   rI   r:   s    r&   rJ   z.PyannoteAudioPretrainedSpeakerEmbedding.metric  rK   r'   c                 v   t        j                         5  dt        d| j                  z        }}||z   dz  }|dz   |k  rS	 | j	                  t        j
                  dd|      j                  | j                              }|}||z   dz  }|dz   |k  rSd d d        |S # t        $ r |}Y (w xY w# 1 sw Y   S xY wr   )	r   r   rO   r5   r!   r   r#   r   	Exceptionr   s        r&   rT   z7PyannoteAudioPretrainedSpeakerEmbedding.min_num_samples  s    !!#eC$*:*:$:;5Eem)F!)e##EKK1f$=$@$@$MNA"E  %-A- !)e# $  ! #"E# $ s0   +B.ABB.B+(B.*B++B..B8rU   rV   c                    t        j                         5  |+| j                  |j                  | j                              }nwt        j                         5  t        j                  d       | j                  |j                  | j                        |j                  | j                              }d d d        d d d        j                         j                         S # 1 sw Y   /xY w# 1 sw Y   3xY w)Nignoreweights)
r   r   r!   r#   r   r   catch_warningssimplefilterr   rl   )r$   rU   rV   rD   s       r&   ry   z0PyannoteAudioPretrainedSpeakerEmbedding.__call__  s     !!#}![[dkk)BC
,,.))(3!%!T[[1588DKK;P "- "J /	 $ ~~%%'' /.	 $#s%   AC,AC 2C, C)	%C,,C5pyannote/embeddingNNNrz   )r0   r{   r|   r   r   r   r   r   r   r   r   r   r#   r   r}   r5   rE   r~   rJ   rT   r   ri   r   ry   r   r   s   @r&   r   r   n  s   < $8)-#'-1$ $ &$ T4Z 	$
 tT)*$  -S - - %3 % %        HL((.6u||.D(	(r'   r   r   r   r   r   c                 4   t        | t              rd| v rt        | |||      S t        | t              rd| v rt        | |||      S t        | t              rd| v rt	        | |      S t        | t              rd| v rt        | |||      S t        | |||      S )a~  Pretrained speaker embedding

    Parameters
    ----------
    embedding : Text
        Can be a SpeechBrain (e.g. "speechbrain/spkrec-ecapa-voxceleb")
        or a pyannote.audio model.
    device : torch.device, optional
        Device
    token : str or bool, optional
        Huggingface token to be used for downloading from Huggingface hub.
    cache_dir: Path or str, optional
        Path to the folder where files downloaded from Huggingface hub are stored.

    Usage
    -----
    >>> get_embedding = PretrainedSpeakerEmbedding("pyannote/embedding")
    >>> get_embedding = PretrainedSpeakerEmbedding("speechbrain/spkrec-ecapa-voxceleb")
    >>> get_embedding = PretrainedSpeakerEmbedding("nvidia/speakerverification_en_titanet_large")
    >>> assert waveforms.ndim == 3
    >>> batch_size, num_channels, num_samples = waveforms.shape
    >>> assert num_channels == 1
    >>> embeddings = get_embedding(waveforms)
    >>> assert embeddings.ndim == 2
    >>> assert embeddings.shape[0] == batch_size

    >>> assert masks.ndim == 1
    >>> assert masks.shape[0] == batch_size
    >>> embeddings = get_embedding(waveforms, masks=masks)
    pyannote)r   r   r   speechbrainnvidia)r   	wespeaker)r-   r~   r   r   r   r   )r   r   r   r   s       r&   PretrainedSpeakerEmbeddingr    s    J )S!jI&=6fEY
 	
 
Is	#(B4fEY
 	
 
Is	#I(=-iGG	Is	#y(@6fEY
 	
 7fEY
 	
r'   c                   ~     e Zd ZdZ	 	 	 	 ddedee   deedf   deeedf   f fdZ	de
d	ej                  fd
Z xZS )SpeakerEmbeddinga  Speaker embedding pipeline

    This pipeline assumes that each file contains exactly one speaker
    and extracts one single embedding from the whole file.

    Parameters
    ----------
    embedding : Model, str, or dict, optional
        Pretrained embedding model. Defaults to "pyannote/embedding".
        See pyannote.audio.pipelines.utils.get_model for supported format.
    segmentation : Model, str, or dict, optional
        Pretrained segmentation (or voice activity detection) model.
        See pyannote.audio.pipelines.utils.get_model for supported format.
        Defaults to no voice activity detection.
    token : str or bool, optional
        Huggingface token to be used for downloading from Huggingface hub.
    cache_dir: Path or str, optional
        Path to the folder where files downloaded from Huggingface hub are stored.

    Usage
    -----
    >>> from pyannote.audio.pipelines import SpeakerEmbedding
    >>> pipeline = SpeakerEmbedding()
    >>> emb1 = pipeline("speaker1.wav")
    >>> emb2 = pipeline("speaker2.wav")
    >>> from scipy.spatial.distance import cdist
    >>> distance = cdist(emb1, emb2, metric="cosine")[0,0]
    Nr   segmentationr   r   c                     t         |           || _        || _        t	        |||      | _        | j                  ,t	        | j                  ||      }t        |d       | _        y y )Nr   c                 2    t        j                  | dd      S )NT)axiskeepdims)ri   rh   )scoress    r&   <lambda>z+SpeakerEmbedding.__init__.<locals>.<lambda>A  s    BFFd5r'   )pre_aggregation_hook)r   r   r   r  r   embedding_model_r   _segmentation)r$   r   r  r   r   segmentation_modelr%   s         r&   r   zSpeakerEmbedding.__init__+  su     	"('0Ui(
 ((1!!)) "+"&"D	 )r'   filer3   c                 "   | j                   j                  }| j                   j                  |      d   d    j                  |      }| j                  d }nb| j                  |      j                  }d|t        j                  |      <   t        j                  |dz        d d d df   j                  |      }t        j                         5  | j                  ||      j                         j                         cd d d        S # 1 sw Y   y xY w)Nr   r      r   )r  r   r   r#   r  r  datari   r   r   
from_numpyno_gradr   rl   )r$   r  r   ru   r   s        r&   applyzSpeakerEmbedding.applyF  s    &&-- ((..t4Q7=@@H$G ((.33G),GBHHW%&&&wz24A:>AA&IG ]]_((7(CGGIOOQ __s   /DDr   )r0   r{   r|   r   r   r   r   r   r   r   r   ri   r   r  r   r   s   @r&   r  r    su    > $804#'-1  }- T4Z 	
 tT)*6R) R

 Rr'   r  protocolsubsetr  c                 f   dd l }ddlm}m} ddlm} ddlm} ddlm}	 t        ||      }
 || d |       i      } g g }}t               } t        | | d	             }t         |	|            D ]m  \  }}|d
   d   }||vr |
|      ||<   |d   d   }||vr |
|      ||<   |j                   |||   ||   d      d   d          |j                  |d          o  ||t        j                  |      d      \  }}}} |j                   | j"                   d| d| d| dd|z  dd
       y )Nr   )
FileFinderget_protocol)	det_curve)cdist)tqdm)r   r  r   )preprocessors_trialfile1file2rH   )rJ   	referenceT)	distancesz | z	 | EER = d   z.3f%)typerpyannote.databaser  r  &pyannote.metrics.binary_classificationr  scipy.spatial.distancer  r  r  dictgetattrr   appendri   arrayechoname)r  r  r   r  r'  r  r  r  r  r  pipeliney_truey_predembtrialsttrialaudio1audio2rC   eers                        r&   mainr;  Z  sO    :@,),OHHWjl4KLHFF
&C1WX&013Fd6l+5w("6*CKw("6*CKeCKVXFqI!LMeK() , VRXXf%5FLAq!SEJJ==/VHC	{#l^9SSVYWZO[\]r'   __main__)NNN)z&VoxCeleb.SpeakerVerification.VoxCeleb1testr   N)<r   	functoolsr   pathlibr   typingr   r   r   rl   ri   r   torch.nn.functionalnn
functionalrc   torchaudio.compliance.kaldi
compliancer   huggingface_hubr   huggingface_hub.utilsr	   torch.nn.utils.rnnr
   pyannote.audior   r   r   pyannote.audio.core.inferencer   pyannote.audio.core.ior   pyannote.audio.pipelines.utilsr   r   speechbrain.inferencer   r   r   r   nemo.collections.asr.modelsr   r   r   onnxruntimer   r   r   r   r   r   r   r  r  r~   r;  r0   r'  r   rI   r'   r&   <module>rP     s  .  %  ( (     + + + 9 + 5 5 7 , C%X# 
F] FRuM upim iX^(m ^(F &*#)-	;
;
U\\";
 t;
 T4%&	;
|JRx JR\ =)"&	%%% % 3-	%P zEIIdO k  %$%    s6   &D /D 8D( DDD%$D%(D21D2