
     i90                         d dl mZ d dlmZmZmZ d dl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 d dlmZ d	ej,                  d
ej,                  fdZ G d de      Z G d de      Z G d de      Zy)    )singledispatchmethod)DictListOptionalN)
AnnotationSegmentSlidingWindowSlidingWindowFeatureTimeline)
BaseMetric) DetectionPrecisionRecallFMeasure)DiarizationErrorRate)	permutate	reference
hypothesisc                    | j                  t        j                        } |j                  t        j                        }t        | t        j                     |      \  \  }}dt        j
                  |       z  }t        j
                  |d      t        j
                  | d      z
  }t        j                  d|      }t        j                  d|       }t        j
                  || k7  |z  d      |z
  }t        j
                  |      }t        j
                  |      }t        j
                  |      }||z   |z   |z  }|||||dfS )a  Discrete diarization error rate

    Parameters
    ----------
    reference : (num_frames, num_speakers) np.ndarray
        Discretized reference diarization.
        reference[f, s] = 1 if sth speaker is active at frame f, 0 otherwise
    hypothesis : (num_frames, num_speakers) np.ndarray
        Discretized hypothesized diarization.
       hypothesis[f, s] = 1 if sth speaker is active at frame f, 0 otherwise

    Returns
    -------
    der : float
        (false_alarm + missed_detection + confusion) / total
    components : dict
        Diarization error rate components, in number of frames.
        Keys are "false alarm", "missed detection", "confusion", and "total".
          ?   )axisr   )false alarmmissed detection	confusiontotal)astypenphalfr   newaxissummaximum)	r   r   _r   detection_errorfalse_alarmmissed_detectionr   ders	            p/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/audio/utils/metric.pydiscrete_diarization_error_rater&   )   s'   *   )I""277+J !2::!6
CMZ1 "&&##E ffZa0266)!3LLO**Q0Kzz!o%56 
i/:=AFTI&&%Kvv./y!I))I5
>C 	& 0"		
     c            	          e Zd ZdZed        Zed        Z	 ddee   fdZ	e
	 ddee   fd       Zej                  	 ddej                  d	ej                  dee   fd
       Zej                  	 dded	edee   fd       Zd Zy)DiscreteDiarizationErrorRatez9Compute diarization error rate on discretized annotationsc                      y)Nzdiscrete diarization error rate clss    r%   metric_namez(DiscreteDiarizationErrorRate.metric_namec   s    0r'   c                 
    g dS N)r   r   r   r   r+   r,   s    r%   metric_componentsz.DiscreteDiarizationErrorRate.metric_componentsg       HHr'   Nuemc                 *    | j                  |||      S )Nr3   )compute_components_helper)selfr   r   r3   s       r%   compute_componentsz/DiscreteDiarizationErrorRate.compute_componentsk   s     --j)-MMr'   c                 L    |j                   j                  }t        d| d      )NzProviding hypothesis as z instances is not supported.)	__class____name__NotImplementedError)r7   r   r   r3   klasss        r%   r6   z6DiscreteDiarizationErrorRate.compute_components_helpers   s/     $$--!&ug-IJ
 	
r'   r   r   c                    |j                   dk7  rt        d      |t        d      |j                  \  }}|j                   dk7  rt        d      |j                  \  }}||k7  rt        d      ||kD  rt	        j
                  |dd||z
  ff      }n"||kD  rt	        j
                  |dd||z
  ff      }t        ||      d   S )	N   z>Only (num_frames, num_speakers)-shaped reference is supported.z)`uem` is not supported with numpy arrays.z?Only (num_frames, num_speakers)-shaped hypothesis is supported.z=reference and hypothesis must have the same number of frames.)r   r   r   r   )ndimr<   
ValueErrorshaper   padr&   )r7   r   r   r3   ref_num_framesref_num_speakershyp_num_frameshyp_num_speakerss           r%   der_from_ndarrayz-DiscreteDiarizationErrorRate.der_from_ndarray|   s     >>Q%P  ?HII+4??((??a%Q  ,6+;+;((^+O  ..FQ(8;K(K$LMI  00Va)9<L)L%MNJ /y*EaHHr'   c                 L   |j                   j                  }|dk  s|dkD  rt        d      |dk(  r|j                  }|j                  }nc|dk(  r^|j                  }|j                   j
                  \  }}	}
t        |d   j                  ||dz
     j                        }|j                  |	z  }|j                        }|dk(  r|&| j                  |j                   |j                         S t        |g      j                  |      st        d      | j                         }|D ]W  }|j!                  |      }|j!                  |      }| j                  ||      }| j"                  D ]  }||xx   ||   z  cc<    Y |S |dk(  r| j                         }|D ]  \  }}||j                  t        |g            s$|j!                  |d	      }t%        	|j
                  d         }| j                  |d | |d |       }| j"                  D ]  }||xx   ||   z  cc<     |S y )
Nr?      ziOnly (num_frames, num_speakers) or (num_chunks, num_frames, num_speakers)-shaped hypothesis is supported.r   r   )
resolutionz)`uem` must fully cover hypothesis extent.center)mode)datar@   r<   extentsliding_windowrB   r   startendduration
discretizer6   r   coversrA   init_componentscropcomponents_min)r7   r   r   r3   r@   supportrK   chunks
num_chunks
num_framesr    
componentssegmenthrsegment_componentnamewindowhypothesis_windowreference_windowcommon_num_frameswindow_componentss                         r%   der_from_swfz)DiscreteDiarizationErrorRate.der_from_swf   sO    ##!8tax%+  19 ''G#22J QY..F(2(=(=%J
AfQioovj1n/E/I/IJG:5J ((Z(H	 19{55jooy~~VVWI&--c2 !LMM--/JOOG,NN7+$($B$B1a$H! ,,Dt$(9$(??$ -	   QY--/J-7)) ?3::hx6H+I#,>>&x>#H $'
4D4J4J14M$N!$($B$B%&8'89$%7&78%!
 !,,Dt$(9$(??$ - .8$ + r'   c                 0    |d   |d   z   |d   z   |d   z  S Nr   r   r   r   r+   r7   r^   s     r%   compute_metricz+DiscreteDiarizationErrorRate.compute_metric   9    }%+,-%& w	  	 r'   N)r;   
__module____qualname____doc__classmethodr.   r1   r   r   r8   r   r6   registerr   ndarrayrH   r
   r   ri   rm   r+   r'   r%   r)   r)   `   s   C1 1 I I #'	N h	N ?C
*28*<
 
 ''
 #'	&IJJ&I ::&I h	&I (&IP ''
 #'	E(E E h	E (EN r'   r)   c                   b     e Zd Zddef fdZed        Zed        Z	 d	dee	   fdZ
d Z xZS )
SlidingDiarizationErrorRaterd   c                 0    t         |           || _        y ro   )super__init__rd   )r7   rd   r:   s     r%   rz   z$SlidingDiarizationErrorRate.__init__   s    r'   c                      y)Nzwindow diarization error rater+   r,   s    r%   r.   z'SlidingDiarizationErrorRate.metric_name   s    .r'   c                 
    g dS r0   r+   r,   s    r%   r1   z-SlidingDiarizationErrorRate.metric_components   r2   r'   r3   c                    |t        d      t               }t        | j                  d| j                  z        } ||      D ]5  } ||j	                  |      |j	                  |      t        |g            }7 |d d  S )Nz9SlidingDiarizationErrorRate expects `uem` to be provided.g      ?)rS   stepr5   )rA   r   r	   rd   rW   r   )r7   r   r   r3   r$   rd   chunkr    s           r%   r8   z.SlidingDiarizationErrorRate.compute_components  s     ;K  #$#:KLC[Eu%zu'=8UGCTA !
 1vr'   c                 0    |d   |d   z   |d   z   |d   z  S rk   r+   rl   s     r%   rm   z*SlidingDiarizationErrorRate.compute_metric  rn   r'   )g      $@ro   )r;   rp   rq   floatrz   rs   r.   r1   r   r   r8   rm   __classcell__r:   s   @r%   rw   rw      sW    u  / / I I #'	 h	. r'   rw   c                        e Zd ZdZd Zed        Z	 	 ddee   de	de	fdZ
 fdZ	 dd	ed
efdZdeee	f   fdZd fd	Zd Z xZS )MacroAverageFMeasurea   Compute macro-average F-measure

    Parameters
    ----------
    collar : float, optional
        Duration (in seconds) of collars removed from evaluation around
        boundaries of reference segments (one half before, one half after).
    beta : float, optional
        When beta > 1, greater importance is given to recall.
        When beta < 1, greater importance is given to precision.
        Defaults to 1.

    See also
    --------
    pyannote.metrics.detection.DetectionPrecisionRecallFMeasure
    c                     | j                   S ro   )classes)r7   s    r%   r1   z&MacroAverageFMeasure.metric_components3  s    ||r'   c                      y)NzMacro F-measurer+   r,   s    r%   r.   z MacroAverageFMeasure.metric_name6  s     r'   r   collarbetac           
         | j                         | _        || _        t        | j	                               | _        || _        || _        | j                  D ci c]  }|t        d||d| c}| _	        | j                          y c c}w )N)r   r   r+   )r.   metric_name_r   setr1   rX   r   r   r   _sub_metricsreset)r7   r   r   r   kwargslabels         r%   rz   zMacroAverageFMeasure.__init__:  s     !,,.t5578	 J
% 3W6WPVWW%J

 	

J
s   Bc                     t         |           | j                  j                         D ]  }|j                           y ro   )ry   r   r   values)r7   
sub_metricr:   s     r%   r   zMacroAverageFMeasure.resetP  s1    ++224J 5r'   r   r   c                     | j                         }| j                  j                         D ]4  \  }} |d|j                  |g      |j                  |g      |d|||<   6 |S )N)r   r   r3   r+   )rV   r   itemssubset)r7   r   r   r3   r   detailsr   r   s           r%   r8   z'MacroAverageFMeasure.compute_componentsU  sw     &&(!%!2!2!8!8!:E:' #**E73%,,eW5 	GEN "; r'   detailc                 Z    t        j                  t        |j                                     S ro   )r   meanlistr   )r7   r   s     r%   rm   z#MacroAverageFMeasure.compute_metricc  s    wwtFMMO,--r'   c                     t         |   d      }| j                  j                         D ]   \  }}t	        |      |j
                  d   |<   " |rt        |j                  dddd              |S )NF)displayTOTALTrightc                 $    dj                  |       S )Nz{0:.2f})format)fs    r%   <lambda>z-MacroAverageFMeasure.report.<locals>.<lambda>r  s    9+;+;A+>r'   )indexsparsifyjustifyfloat_format)ry   reportr   r   abslocprint	to_string)r7   r   dfr   r   r:   s        r%   r   zMacroAverageFMeasure.reportf  sx    W^E^*!%!2!2!8!8!:E:%(_BFF7OE" "; "#!>	   	r'   c                     t        j                  | j                  j                         D cg c]  }t	        |       c}      S c c}w ro   )r   r   r   r   r   )r7   r   s     r%   __abs__zMacroAverageFMeasure.__abs__x  s8    ww$:K:K:R:R:TU:TJJ:TUVVUs   A)g        r   ro   )F)r;   rp   rq   rr   r1   rs   r.   r   strr   rz   r   r   r8   r   rm   r   r   r   r   s   @r%   r   r   !  s    " ! ! 	c  	, BF#1;.T#u*%5 .$Wr'   r   )	functoolsr   typingr   r   r   numpyr   pyannote.corer   r   r	   r
   r   pyannote.metrics.baser   pyannote.metrics.detectionr   pyannote.metrics.diarizationr    pyannote.audio.utils.permutationr   ru   r&   r)   rw   r   r+   r'   r%   <module>r      ss   . + ' '   - G = 64rzz 4rzz 4nR : R j) * ) XXW: XWr'   