
     iFu                     d   d Z ddlZddlZddlZddlZddlZddlmZ ddlm	Z	m
Z
mZmZmZmZ ddlmZ ddlZddlZddlmZ ddlmZmZmZ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$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dl0m1Z1m2Z2 dde3fdZ4e G d d             Z5 G d de$e      Z6y)zSpeaker diarization pipelines    N)Path)CallableMappingOptionalTextUnionAny)	dataclass)	rearrange)Audio	InferenceModelPipeline)	AudioFile)
Clustering)PretrainedSpeakerEmbedding)PipelineModelPipelinePLDASpeakerDiarizationMixin	get_modelget_plda)set_num_speakers)binarize)
AnnotationSlidingWindowFeature)GreedyDiarizationErrorRate)	ParamDictUniform
batch_sizec                 J    t        |       g|z  }t        j                  |d|iS )zBatchify iterable	fillvalue)iter	itertoolszip_longest)iterabler   r!   argss       /Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/pipelines/speaker_diarization.pybatchifyr(   8   s+     Nj(D  $<)<<    c                   `    e Zd ZU eed<   eed<   dZej                  dz  ed<   dee	e
f   fdZy)DiarizeOutputspeaker_diarizationexclusive_speaker_diarizationNspeaker_embeddingsreturnc                    g }| j                   j                  d      D ]C  \  }}}|j                  t        |j                  d      t        |j
                  d      |d       E g }| j                  j                  d      D ]C  \  }}}|j                  t        |j                  d      t        |j
                  d      |d       E ||dS )a  Serialize diarization output

        Example
        -------
        {
            'diarization': [{
                'start': 6.665,
                'end': 7.165,
                'speaker': 'SPEAKER_00'},
                ...],
            'exclusive_diarization': [{
                'start': 6.665,
                'end': 7.165,
                'speaker': 'SPEAKER_00'},
                ...],
        }
        T)yield_label   )startendspeaker)diarizationexclusive_diarization)r,   
itertracksappendroundr3   r4   r-   )selfr6   speech_turn_r5   r7   s         r'   	serializezDiarizeOutput.serializeN   s    & '+'?'?'J'J (K (
#KG ";#4#4a8 !4&(
 !#'+'I'I'T'T (U (
#KG "((";#4#4a8 !4&(
 '%:
 	
r)   )__name__
__module____qualname__r   __annotations__r.   npndarraydictstrr	   r>    r)   r'   r+   r+   ?   s<     $# $.-
 -1

T)0.
4S> .
r)   r+   c                       e Zd ZdZddddddddddddd	d
d
dddfdededededededede	de	de
e   deedf   deeedf   f fdZede	fd       Zej$                  de	fd       Zd Zd Zed        Zd,defdZ	 	 d-deded e
e   fd!Zd"ed#ej6                  d$edefd%Z	 	 	 	 d.d&ed'e
e	   d(e
e	   d)e
e	   d e
e   deez  fd*Z de!fd+Z" xZ#S )/SpeakerDiarizationa  Speaker diarization pipeline

    Parameters
    ----------
    legacy : bool, optional
        Return only the diarization output. Defaults to return the full output
        with diarization, exclusive diarization, and speaker embeddings.
    segmentation : Model, str, or dict, optional
        Pretrained segmentation model. 
        See pyannote.audio.pipelines.utils.get_model for supported format.
    segmentation_step: float, optional
        The segmentation model is applied on a window sliding over the whole audio file.
        `segmentation_step` controls the step of this window, provided as a ratio of its
        duration. Defaults to 0.1 (i.e. 90% overlap between two consecuive windows).
    embedding : Model, str, or dict, optional
        Pretrained embedding model. 
        See pyannote.audio.pipelines.utils.get_model for supported format.
    embedding_exclude_overlap : bool, optional
        Exclude overlapping speech regions when extracting embeddings.
        Defaults (False) to use the whole speech.
    plda : PLDA, str, or dict, optional
        Pretrained PLDA.
        See pyannote.audio.pipelines.utils.get_plda for supported format.
    clustering : str, optional
        Clustering algorithm. See pyannote.audio.pipelines.clustering.Clustering
        for available options. 
    segmentation_batch_size : int, optional
        Batch size used for speaker segmentation. Defaults to 1.
    embedding_batch_size : int, optional
        Batch size used for speaker embedding. Defaults to 1.
    der_variant : dict, optional
        Optimize for a variant of diarization error rate.
        Defaults to {"collar": 0.0, "skip_overlap": False}. This is used in `get_metric`
        when instantiating the metric: GreedyDiarizationErrorRate(**der_variant).
    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
    -----
    # process audio file
    >>> output = pipeline("/path/to/audio.wav")

    # print diarization
    >>> assert isinstance(output.speaker_diarization, pyannote.core.Annotation)
    >>> for turn, speaker in output.speaker_diarization:
    ...     print(f"start={turn.start:.1f}s stop={turn.end:.1f}s speaker_{speaker}")
    
    # get one speaker embedding per speaker
    >>> assert isinstance(output.speaker_embeddings, np.ndarray)
    >>> for s, speaker in enumerate(output.speaker_diarization.labels()):
    ...     # output.speaker_embeddings[s] is the embedding of speaker `speaker`

    # exclusive diarization is the same as diarization except 
    # that it does not contain overlapping speech segments
    >>> assert isinstance(output.exclusive_speaker_diarization, pyannote.core.Annotation)

    # force exactly 4 speakers
    >>> output = pipeline("/path/to/audio.wav", num_speakers=4)

    # force between 2 and 10 speakers
    >>> output = pipeline("/path/to/audio.wav", min_speakers=2, max_speakers=10)
    Fz(pyannote/speaker-diarization-community-1segmentation)
checkpoint	subfolder皙?	embeddingpldaVBxClustering   Nlegacysegmentation_stepembedding_exclude_overlap
clusteringembedding_batch_sizesegmentation_batch_sizeder_varianttoken	cache_dirc           	      n   t         |           || _        || _        t	        |||      }|| _        || _        || _        || _        || _	        t        |||      | _        || _        |
xs ddd| _        |j                  j                  }t!        ||| j
                  |z  d|	      | _        | j"                  j$                  j                  j&                  rt)        t+        dd            | _        n&t)        t+        d	d
      t+        dd            | _        | j                  dk(  rd}nYt/        | j                  ||      | _        t3        | j0                  j4                  d      | _        | j0                  j8                  }	 t:        |   }| j                  dk(  r#|jG                  | j                  |      | _$        n|jG                  |      | _$        | jH                  jJ                  | _&        y # t<        $ r6 t?        ddjA                  tC        t:        jD                               d      w xY w)N)rY   rZ           F)collarskip_overlapT)durationstepskip_aggregationr         ?)min_duration_offrM   g?)	thresholdrc   OracleClusteringnot_applicabledownmix)sample_ratemonozclustering must be one of [, ]rP   )metric)'super__init__rR   segmentation_modelr   rS   rN   rV   rT   rO   r   _plda
klusteringrX   specificationsr_   r   _segmentationmodelpowersetr   r   rJ   r   
_embeddingr   rh   _audiorl   r   KeyError
ValueErrorjoinlist__members__valuerU   expects_num_clusters_expects_num_speakers)r;   rR   rJ   rS   rN   rT   rO   rU   rV   rW   rX   rY   rZ   rt   segmentation_durationrl   
Klustering	__class__s                    r'   rn   zSpeakerDiarization.__init__   s   0 	". UiP!2"$8!)B&	d%9E
$&PS%*P % 4 4 = =&*''*??!.
 ##22;; )!(c!2!D
 !*!#s+!(c!2!D
 ??00%F 9eyDO  DOO,G,GiXDK__++F	#J/J ??o-(..tzz&.IDO(..f.=DO%)__%I%I"  	-diiZ=S=S8T.U-VVWX 	s   	G5 5?H4r/   c                 .    | j                   j                  S Nrs   r   r;   s    r'   rW   z*SpeakerDiarization.segmentation_batch_size  s    !!,,,r)   r   c                 &    || j                   _        y r   r   )r;   r   s     r'   rW   z*SpeakerDiarization.segmentation_batch_size  s    (2%r)   c                     ddidddddS )Nrc   r\   g333333?gQ?g?)rd   FaFb)rJ   rU   rG   r   s    r'   default_parametersz%SpeakerDiarization.default_parameters!  s    /5(+4sC
 	
r)   c              #   ,   K   d}	 d|d |dz  }w)Nr   SPEAKER_02drQ   rG   )r;   r5   s     r'   classeszSpeakerDiarization.classes'  s+     WSM**qLG s   c                      y)Nztraining_cache/segmentationrG   r   s    r'   CACHED_SEGMENTATIONz&SpeakerDiarization.CACHED_SEGMENTATION-  s    ,r)   c                     |t        j                  |dd      }| j                  rC| j                  |v r|| j                     }|S | j	                  ||      }||| j                  <   |S | j	                  ||      }|S )zApply segmentation model

        Parameter
        ---------
        file : AudioFile
        hook : Optional[Callable]

        Returns
        -------
        segmentations : (num_chunks, num_frames, num_speakers) SlidingWindowFeature
        NrJ   hook)	functoolspartialtrainingr   rs   )r;   filer   segmentationss       r'   get_segmentationsz$SpeakerDiarization.get_segmentations1  s     $$T>4@D==''4/ $T%=%= >  !% 2 24d 2 C1>T--.  372D2DTPT2D2UMr)   binary_segmentationsexclude_overlapr   c                      j                   rij                  dt                     }d|v rK j                  j                  j
                  j                  s|d    j                  j                  k(  r|d   S j                  j                  }j                  j                  \  }}}	|r j                  j                  }
| j                  j                  z  }t!        j"                  ||
z  |z        dt%        j&                  j                  dd      dk  z  }t)        j                  |z  j                        n"dt)        j                  j                         fd	}t+         |        j,                  d
      }t!        j"                  ||	z   j,                  z        }g }| |dd|d       t/        |d      D ]x  \  }}t1        t3        d |       \  }}t5        j6                  |      }t5        j6                  |      } j                  ||      }|j9                  |       |m |d|||       z t%        j6                  |      }t;        |d|      } j                   rO j                  j                  j
                  j                  r	d|id<   |S  j                  j                  |dd<   |S )a  Extract embeddings for each (chunk, speaker) pair

        Parameters
        ----------
        file : AudioFile
        binary_segmentations : (num_chunks, num_frames, num_speakers) SlidingWindowFeature
            Binarized segmentation.
        exclude_overlap : bool, optional
            Exclude overlapping speech regions when extracting embeddings.
            In case non-overlapping speech is too short, use the whole speech.
        hook: Optional[Callable]
            Called during embeddings after every batch to report the progress

        Returns
        -------
        embeddings : (num_chunks, num_speakers, dimension) array
        ztraining_cache/embeddings
embeddingssegmentation.thresholdrb      T)axiskeepdimsc               3     K   t        	      D ]  \  \  } }\  }}j                  j                  
| d      \  }}t        j                  |d      j                  t        j                        }t        j                  |d      j                  t        j                        }t        |j                  |j                        D ]A  \  }}t        j                  |      kD  r|}n|}|d    t        j                  |      d    f C  y w)Npad)moder\   )nan)ziprw   croprC   
nan_to_numastypefloat32Tsumtorch
from_numpy)chunkmasksr=   clean_maskswaveformmask
clean_mask	used_maskr   clean_segmentationsr   min_num_framesr;   s           r'   iter_waveform_and_maskzASpeakerDiarization.get_embeddings.<locals>.iter_waveform_and_mask  s     47$&950 0K #kk.. / ! e5<<RZZH mmKSAHHT(+EGG[]](C$D* vvj)N:$.	$(	"4.%*:*:9*Ed*KKK )D#5s   D	D)NN)r   r!   Nr   )total	completedrQ   c                     | d   d uS )Nr   rG   )bs    r'   <lambda>z3SpeakerDiarization.get_embeddings.<locals>.<lambda>  s    QqT5Er)   )r   z(c s) d -> c s d)c)r   r   )r   getrE   rs   rt   rr   ru   rJ   rd   sliding_windowr_   datashaperv   min_num_samplesrh   mathceilrC   r   r   r(   rV   	enumerater   filterr   vstackr9   r   )r;   r   r   r   r   cacher_   
num_chunks
num_framesnum_speakersr   num_samplesclean_framesr   batchesbatch_countembedding_batchesibatch	waveformsr   waveform_batch
mask_batchembedding_batchr   r   r   s   ```                      @@r'   get_embeddingsz!SpeakerDiarization.get_embeddingsL  s   : == HH8$&AE%""((77@@23t7H7H7R7RR\**'66??/C/H/H/N/N,
J #oo==O #T__%@%@@K!YYzO'Ck'QRN +00q4H1LL #7$))L8$33#  N"6$))+?+N+N#	L 	L< "$00"
 ii
\ 9D<U<U UVt;!D!'1-HAu"F+Eu$MNIu"\\)4N e,J +///j +: +O
 $$_5\?+QRS# .& II&7802D
S
 ==!!''66?? *501 	 /3.?.?.I.I",501
 r)   r   hard_clusterscountc                    |j                   j                  \  }}}t        j                  |      dz   }t        j                  t        j
                  |||f      z  }t        t        ||            D ]T  \  }	\  }
\  }}t        j                  |
      D ]1  }|dk(  r	t        j                  |dd|
|k(  f   d      ||	dd|f<   3 V t        ||j                        }| j                  ||      S )a;  Build final discrete diarization out of clustered segmentation

        Parameters
        ----------
        segmentations : (num_chunks, num_frames, num_speakers) SlidingWindowFeature
            Raw speaker segmentation.
        hard_clusters : (num_chunks, num_speakers) array
            Output of clustering step.
        count : (total_num_frames, 1) SlidingWindowFeature
            Instantaneous number of active speakers.

        Returns
        -------
        discrete_diarization : SlidingWindowFeature
            Discrete (0s and 1s) diarization.
        rQ   Nr   )r   r   rC   maxr   zerosr   r   uniquer   r   to_diarization)r;   r   r   r   r   r   local_num_speakersnum_clustersclustered_segmentationsr   clusterr   rJ   ks                 r'   reconstructzSpeakerDiarization.reconstruct  s    . 6C5G5G5M5M2
J 2vvm,q0"$&&288\2,
 #
 4=}-4
/A/.%
 YYw'7 4666 GqL14'1a0 (4
 #7#]%A%A#
 ""#:EBBr)   r   r   min_speakersmax_speakersc                    t        |      dkD  r>t        j                  ddj                  t	        |j                                             | j                  ||      }t        |||      \  }}}| j                  rL|Jt        |t              r!d|v rt        |d   j                               }nt        d| j                   d      | j                  ||      } |d	|       |j                  j                   \  }}	}
| j"                  j$                  j&                  j(                  r|}n"t+        || j,                  j.                  d
      }| j1                  || j"                  j$                  j2                  d      } |d|       t5        j6                  |j                        dk(  rkt9        t;        |d         t;        |d         t5        j<                  d| j>                  j@                  f            }| jB                  r|jD                  S |S | jG                  ||| jH                  |      } |d|       | jK                  ||||||| j"                  j$                  j2                        \  }}}t5        jL                  |      dz   }||k  s||kD  r;t        j                  tO        jP                  d| d|d    d| d| d| d             t5        jR                  |j                  |      jU                  t4        jV                        |_        t5        jX                  |j                  d      dk(  }d||<   | j[                  |||      } |d|       | j]                  |d| j,                  j^                        }|d   |_0        t5        jR                  |j                  d      jU                  t4        jV                        |_        | j[                  |||      }| j]                  |d| j,                  j^                        }|d   |_0        d|v rN|d   rI| jc                  |d   |d !      \  }}|j                         D ci c]  }||je                  ||       }}n;tg        |j                         | ji                               D ci c]  \  }}||
 }}}|jk                  |"      }|jk                  |"      }|(t9        |||      }| jB                  r|jD                  S |S t        |j                               |j                   d   kD  rAt5        jl                  |dt        |j                               |j                   d   z
  fd#f      }|jo                         D ci c]  \  }}||
 }}}||j                         D cg c]  }||   	 c}   }t9        |||      }| jB                  r|jD                  S |S c c}w c c}}w c c}}w c c}w )$aC  Apply speaker diarization

        Parameters
        ----------
        file : AudioFile
            Processed file.
        num_speakers : int, optional
            Number of speakers, when known.
        min_speakers : int, optional
            Minimum number of speakers. Has no effect when `num_speakers` is provided.
        max_speakers : int, optional
            Maximum number of speakers. Has no effect when `num_speakers` is provided.
        hook : callable, optional
            Callback called after each major steps of the pipeline as follows:
                hook(step_name,      # human-readable name of current step
                     step_artefact,  # artifact generated by current step
                     file=file)      # file being processed
            Time-consuming steps call `hook` multiple times with the same `step_name`
            and additional `completed` and `total` keyword arguments usable to track
            progress of current step.

        Returns
        -------
        output : DiarizeOutput (or Annotation if `self.legacy` is True)
        r   z'Ignoring unexpected keyword arguments: rj   r   )r   r   r   
annotationz)num_speakers must be provided when using z clusteringrJ   F)onsetinitial_state)r\   r\   )warm_upspeaker_countingr\   uri)r   )r,   r-   r.   )r   r   r   )r   r   r   min_clustersmax_clustersr   framesrQ   z2
                The detected number of speakers (z) for z. is outside
                the given bounds [zS]. This can happen if the
                given audio file is too short to contain zh or more speakers.
                Try to lower the desired minimal number of speakers.
                r   r   discrete_diarization)min_duration_onrc   T)return_mapping)mapping)r   r   )8lenwarningswarnrz   r{   keys
setup_hookr   r   
isinstancer   labelsry   rq   r   r   r   rs   rt   rr   ru   r   rJ   rd   speaker_countreceptive_fieldrC   nanmaxr+   r   r   rv   	dimensionrR   r,   r   rT   rU   r   textwrapdedentminimumr   int8r   r   to_annotationrc   r   optimal_mappingr   r   r   rename_labelsr   items)r;   r   r   r   r   r   kwargsr   r   r   r   binarized_segmentationsr   outputr   r   r=   	centroidsnum_different_speakersinactive_speakersr   r6   exclusive_discrete_diarizationr7   r   keylabelexpected_labelindexinverse_mappings                                 r'   applyzSpeakerDiarization.apply  s   H v;?MM9$))DDW:X9YZ
 t$/3C%%%4
0lL %%,*>$(\T-A"4#5#<#<#>? !??PP[\  ..t$.?^]+5B5G5G5M5M2
J 2 ##22;;&3#<D''11#=# ""#$$44 # 

 	'
 99UZZ C'"$.4;$?.8T%[.I#%88Q0I0I,J#KF {{111M((# ::	 ) 

 	\:& '+oo!1%%%%%++;; '6 '
#q) "$!6!: #\1%4MM22H1IPTUZP[} ]##/.<. A::F H	 ZZ

L9@@I

 FF#:#?#?aHAM ,.'(  $// 

 	#%9:(( !..?? ) 

 u+ ZZ

A.55bgg>
)-)9)9*
&
 !% 2 2*!..?? !3 !

 %)K!
 4D$6 --\"K . JAw >I=O=O=QR=QcsGKKS11=QGR .11C1C1Et||~-V-V)E> ~%-V  
 "///@ 5 C CG C T "$/.C#,F
 {{111M {!!#$yq'99QK$6$6$8 9IOOA<N NOQWXI =DMMOLOLE55%<OL1<1C1C1EF1E_U#1EF
	  +*?(
 ;;---o S
H MFs   W
W.WWc                 ,    t        di | j                  S )NrG   )r   rX   r   s    r'   
get_metriczSpeakerDiarization.get_metric  s    )=D,<,<==r)   r   )FN)NNNN)$r?   r@   rA   __doc__boolr   floatr   rF   intr   rE   r   r   r   rn   propertyrW   setterr   r   r   r   r   r   r   rC   rD   r   r   r+   r   r  r   r  __classcell__)r   s   @r'   rI   rI      sE   ?F D''
 $'D$$
 +0D
 *$%'(&*#'-1-VJVJ $VJ !VJ !VJ $(VJ VJ" #VJ$ "%VJ& "%'VJ( d^)VJ* T4Z +VJ, tT)*-VJp - - - ##3# 3 $3
 - -4H > !&#'R 3R 	R
 x Rh0C+0C zz0C $	0C
 
0Cj '+&*&*#'~~ sm~ sm	~
 sm~ x ~ 
	#~@>6 >r)   rI   )    N)7r  r   r#   r   r  r   pathlibr   typingr   r   r   r   r   r	   dataclassesr
   numpyrC   r   einopsr   pyannote.audior   r   r   r   pyannote.audio.core.ior   #pyannote.audio.pipelines.clusteringr   -pyannote.audio.pipelines.speaker_verificationr   pyannote.audio.pipelines.utilsr   r   r   r   r   *pyannote.audio.pipelines.utils.diarizationr   pyannote.audio.utils.signalr   pyannote.corer   r   pyannote.metrics.diarizationr   pyannote.pipeline.parameterr   r   r  r(   r+   rI   rG   r)   r'   <module>r3     s   0 $       @ @ !    < < , : T  H 0 : C :=3 = <
 <
 <
~T
>0( T
>r)   