
      i                     ^   d dl 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 d dlmZ d dlmZ dd	lmZ 	 dd
e	de	dedeej&                  ej&                  ej&                  ef   fdZ	 dd
e	de	dedeej&                  ej&                  ej&                  ef   fdZ G d de      Z G d d      Zy)    )Counter)TupleN)	ArrayLike)BaseEstimator)CalibratedClassifierCV)_CVIterableWrapper   )CalibrationMethody_truescores	distancesreturnc                     |r| }t         j                  j                  | |d      \  }}}d|z
  }|r| }t        j                  ||kD        d   d   }d||dz
     ||   z   ||dz
     z   ||   z   z  }||||fS )a  DET curve

    Parameters
    ----------
    y_true : (n_samples, ) array-like
        Boolean reference.
    scores : (n_samples, ) array-like
        Predicted score.
    distances : boolean, optional
        When True, indicate that `scores` are actually `distances`

    Returns
    -------
    fpr : numpy array
        False alarm rate
    fnr : numpy array
        False rejection rate
    thresholds : numpy array
        Corresponding thresholds
    eer : float
        Equal error rate
    T	pos_labelr	   r   g      ?)sklearnmetrics	roc_curvenpwhere)	r   r   r   fprtpr
thresholdsfnr	eer_indexeers	            {/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/pyannote/metrics/binary_classification.py	det_curver   *   s    4  #??44VVt4TCj
c'C [
 s#A&q)I
IMS^+c)a-.@@3y>QC Z$$    c                     |r| }t         j                  j                  | |d      \  }}}|r| }t         j                  j                  ||d      }||||fS )a  Precision-recall curve

    Parameters
    ----------
    y_true : (n_samples, ) array-like
        Boolean reference.
    scores : (n_samples, ) array-like
        Predicted score.
    distances : boolean, optional
        When True, indicate that `scores` are actually `distances`

    Returns
    -------
    precision : numpy array
        Precision
    recall : numpy array
        Recall
    thresholds : numpy array
        Corresponding thresholds
    auc : float
        Area under curve

    Tr   )reorder)r   r   precision_recall_curveauc)r   r   r   	precisionrecallr   r#   s          r   r"   r"   W   sl    6 $+OO$J$J$ %K %!Ivz  [

//

i

>Cfj#--r   c                   4     e Zd ZdZ fdZd ZdefdZ xZS )_Passthroughz7Dummy binary classifier used by score Calibration classc                 f    t         |           t        j                  ddgt              | _        y )NFT)dtype)super__init__r   arrayboolclasses_)self	__class__s    r   r+   z_Passthrough.__init__   s$    %d;r   c                     | S N )r/   r   r   s      r   fitz_Passthrough.fit   s    r   r   c                     |S )z"Returns the input scores unchangedr3   r/   r   s     r   decision_functionz_Passthrough.decision_function   s    r   )	__name__
__module____qualname____doc__r+   r4   r   r7   __classcell__)r0   s   @r   r'   r'      s    A<	 r   r'   c                   @    e Zd ZdZ	 d
dedefdZdedefdZdefdZ	y	)Calibrationa;  Probability calibration for binary classification tasks

    Parameters
    ----------
    method : {'isotonic', 'sigmoid'}, optional
        See `CalibratedClassifierCV`. Defaults to 'isotonic'.
    equal_priors : bool, optional
        Set to True to force equal priors. Default behavior is to estimate
        priors from the data itself.

    Examples
    --------
    >>> calibration = Calibration()
    >>> calibration.fit(train_score, train_y)
    >>> test_probability = calibration.transform(test_score)

    See also
    --------
    CalibratedClassifierCV

    equal_priorsmethodc                      || _         || _        y r2   )r@   r?   )r/   r?   r@   s      r   r+   zCalibration.__init__   s     (r   r   r   c                 ~   | j                   rt        |      }|d   |d   }}||kD  r
d\  }}||}	}n	d\  }}||}	}t        d||	z  dz         }
t        j                  ||k(        d   }t        j                  ||k(        d   }g }t        |
      D ]L  }t        j                  t        j                  j                  ||	d      |g      }|j                  g |f       N t        |      }nd	}t        t               | j                  |
      | _        | j                  j                  |j!                  dd      |       | S )zTrain calibration

        Parameters
        ----------
        scores : (n_samples, ) array-like
            Uncalibrated scores.
        y_true : (n_samples, ) array-like
            True labels (dtype=bool).
        TF)TF)FT2   r	   r   )sizereplaceprefit)base_estimatorr@   cv)r?   r   minr   r   rangehstackrandomchoiceappendr   r   r'   r@   calibration_r4   reshape)r/   r   r   counterpositivenegativemajorityminority
n_majority
n_minorityn_splitsminority_indexmajority_indexrH   _
test_indexs                   r   r4   zCalibration.fit   sL    foG!(hH("%0"()18J
%0"()18J
2zZ7!;<HXXf&89!<NXXf&89!<NB8_YY		((*U )  '	
 		2z*+ % $B'B B2'>$++"
 	fnnR3V<r   c                 f    | j                   j                  |j                  dd            dddf   S )a#  Calibrate scores into probabilities

        Parameters
        ----------
        scores : (n_samples, ) array-like
            Uncalibrated scores.

        Returns
        -------
        probabilities : (n_samples, ) array-like
            Calibrated scores (i.e. probabilities)
        rI   r	   N)rP   predict_probarQ   r6   s     r   	transformzCalibration.transform   s/       ..v~~b!/DEadKKr   N)Fisotonic)
r8   r9   r:   r;   r-   r
   r+   r   r4   r`   r3   r   r   r>   r>      sE    . GQ) )2C)3) 3Y 3jL	 Lr   r>   )F)collectionsr   typingr   numpyr   sklearn.metricsr   numpy.typingr   sklearn.baser   sklearn.calibrationr   sklearn.model_selection._splitr   typesr
   r-   ndarrayfloatr   r"   r'   r>   r3   r   r   <module>rm      s   :      " & 6 = $ =B*%*%(*%59*%
2::rzz2::u45*%\ =B'.'.('.59'.
2::rzz2::u45'.T= _L _Lr   