Ë
    êÿæiˆA  ã                   ó°   — d dl 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 d dlmZmZ  G d	„ d
«      Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zy)é    N)ÚMapping)Úcheck_consistent_length)Úcheck_matplotlib_support)Ú_get_response_values_binary)Úparse_version)Útype_of_target)Ú_check_pos_label_consistencyÚ_num_samplesc                   óˆ   — e Zd ZdZdddœd„Zeddddœd„«       Zeddddœd	„«       Zed
„ «       Ze	d„ «       Z
e		 	 dd„«       Zy)Ú"_BinaryClassifierCurveDisplayMixinzØMixin class to be used in Displays requiring a binary classifier.

    The aim of this class is to centralize some validations regarding the estimator and
    the target and gather the response of the estimator.
    N)ÚaxÚnamec          	      óÎ   — t        | j                  j                  › d�«       dd lm} |€|j                  «       \  }}|€t        | dt        | dd «      «      }||j                  |fS )Nz.plotr   Úestimator_namer   )r   Ú	__class__Ú__name__Úmatplotlib.pyplotÚpyplotÚsubplotsÚgetattrÚfigure)Úselfr   r   ÚpltÚ_s        úl/Volumes/fast/ai/experiments/voice-extract-mac/.venv/lib/python3.12/site-packages/sklearn/utils/_plotting.pyÚ_validate_plot_paramsz8_BinaryClassifierCurveDisplayMixin._validate_plot_params   sc   € Ü  D§N¡N×$;Ñ$;Ð#<¸EÐ!BÔCÝ'àˆ:Ø—L‘L“N‰EˆAˆrð ˆ<Ü˜4Ð!1´7¸4ÀÈÓ3NÓOˆDØ�2—9‘9˜dÐ"Ð"ó    Úauto)Úresponse_methodÚ	pos_labelr   c                ó”   — t        | j                  › d�«       |€|j                  j                  n|}t        ||||¬«      \  }}|||fS )Nz.from_estimator)r   r    )r   r   r   r   )ÚclsÚ	estimatorÚXÚyr   r    r   Úy_preds           r   Ú!_validate_and_get_response_valueszD_BinaryClassifierCurveDisplayMixin._validate_and_get_response_values$   sX   € ô 	! C§L¡L >°Ð!AÔBà/3¨|ˆy×"Ñ"×+Ò+Àˆä7ØØØ+Øô	
Ñˆ�	ð �y $Ð&Ð&r   )Úsample_weightr    r   c                óÄ   — t        | j                  › d�«       t        |«      dk7  rt        dt        |«      › d�«      ‚t	        |||«       t        ||«      }|�|nd}||fS )Nz.from_predictionsÚbinaryz The target y is not binary. Got ú type of target.Ú
Classifier)r   r   r   Ú
ValueErrorr   r	   )r"   Úy_truer&   r(   r    r   s         r   Ú!_validate_from_predictions_paramszD_BinaryClassifierCurveDisplayMixin._validate_from_predictions_params5   sz   € ô 	! C§L¡L >Ð1BÐ!CÔDä˜&Ó! XÒ-ÜØ2´>À&Ó3IÐ2Jð Kð óð ô
 	  ¨°Ô>Ü0°¸FÓCˆ	àÐ'‰t¨\ˆà˜$ˆÐr   c                ó˜  ‡— t        | j                  › d�«       ddh}t        ˆfd„|D «       «      st        d|› d�«      ‚t	        ‰d   d   d   «      t	        ‰d   d	   d   «      }}t        |«      ||z   k7  rt        d
||z   › dt        |«      › d�«      ‚t        |«      dk7  rt        dt        |«      › d�«      ‚t        |||«       y )Nz.from_cv_resultsr#   Úindicesc              3   ó&   •K  — | ]  }|‰v –— Œ
 y ­w)N© )Ú.0ÚkeyÚ
cv_resultss     €r   Ú	<genexpr>zV_BinaryClassifierCurveDisplayMixin._validate_from_cv_results_params.<locals>.<genexpr>T   s   øè ø€ Ð>±¨�3˜*Ô$±ùs   ƒzB`cv_results` does not contain one of the following required keys: zr. Set explicitly the parameters `return_estimator=True` and `return_indices=True` to the function`cross_validate`.Útrainr   Útestz=`X` does not contain the correct number of samples. Expected z, got Ú.r*   z"The target `y` is not binary. Got r+   )r   r   Úallr-   Úlenr
   r   r   )r"   r6   r$   r%   r(   Úrequired_keysÚ
train_sizeÚ	test_sizes    `      r   Ú _validate_from_cv_results_paramszC_BinaryClassifierCurveDisplayMixin._validate_from_cv_results_paramsH   s  ø€ ô 	! C§L¡L >Ð1AÐ!BÔCà$ iÐ0ˆÜÓ>±Ó>Ô>ÜØTØ �/ð "$ð$óð ô �
˜9Ñ% gÑ.¨qÑ1Ó2Ü�
˜9Ñ% fÑ-¨aÑ0Ó1ð ˆ
ô
 ˜‹?˜j¨9Ñ4Ò4ÜðØ&¨Ñ2Ð3°6¼,Àq»/Ð9JÈ!ðMóð ô
 ˜!Ó Ò(ÜØ4´^ÀAÓ5FÐ4GÐGWÐXóð ô 	   1 mÕ4r   c                 óR   — | �|�|› d|› d| d›d�}|S | �
|› d| d›�}|S |�|}|S d}|S )z;Helper to get legend label using `name` and `legend_metric`Nz (z = ú0.2fÚ)r3   )Úcurve_legend_metricÚ
curve_nameÚlegend_metric_nameÚlabels       r   Ú_get_legend_labelz4_BinaryClassifierCurveDisplayMixin._get_legend_labelm   sz   € ð Ð*¨zÐ/EØ!�l "Ð%7Ð$8¸Ð<OÐPTÐ;UÐUVÐWˆEð ˆð !Ð,Ø)Ð*¨#Ð.AÀ$Ð-GÐHˆEð
 ˆð	 Ð#ØˆEð ˆð ˆEØˆr   c           	      óä  — |r|rt        d«      ‚|rt        j                  dt        «       |}t	        |t
        «      r t        |«      | k7  rt        d| › d|› d�«      ‚t	        |t
        «      r-t        |«      dk7  rt	        |t
        «      st        d| › d�«      ‚t	        |t        «      r|g}t	        |t
        «      rt        |«      dk(  r|| z  }|€d	g| z  n|}t	        |t        «      r|g| z  }n|€i g| z  }|€i }|€i }| dkD  r|j                  |«       g }d
|v rat        j                  |d
   |d   |«      }	|d   �"|d   �|	d	d d|d   d›d�z   }	n|	d|d   d›�z   }	|j                  |	gd	g| dz
  z  z   «       n=t        |d   |«      D ]+  \  }
}|j                  t        j                  |
||«      «       Œ- t        |«      D ��cg c]  \  }}t!        d|i|¥||   «      ‘Œ }}}|S c c}}w )a�  Get validated line kwargs for each curve.

        Parameters
        ----------
        n_curves : int
            Number of curves.

        name : list of str or None
            Name for labeling legend entries.

        legend_metric : dict
            Dictionary with "mean" and "std" keys, or "metric" key of metric
            values for each curve. If None, "label" will not contain metric values.

        legend_metric_name : str
            Name of the summary value provided in `legend_metrics`.

        curve_kwargs : dict or list of dict or None
            Dictionary with keywords passed to the matplotlib's `plot` function
            to draw the individual curves. If a list is provided, the
            parameters are applied to the curves sequentially. If a single
            dictionary is provided, the same parameters are applied to all
            curves.

        default_curve_kwargs : dict, default=None
            Default curve kwargs, to be added to all curves. Individual kwargs
            are over-ridden by `curve_kwargs`, if kwarg also set in `curve_kwargs`.

        default_multi_curve_kwargs : dict, default=None
            Default curve kwargs for multi-curve plots. Individual kwargs
            are over-ridden by `curve_kwargs`, if kwarg also set in `curve_kwargs`.

        **kwargs : dict
            Deprecated. Keyword arguments to be passed to matplotlib's `plot`.
        z­Cannot provide both `curve_kwargs` and `kwargs`. `**kwargs` is deprecated in 1.7 and will be removed in 1.9. Pass all matplotlib arguments to `curve_kwargs` as a dictionary.z}`**kwargs` is deprecated and will be removed in 1.9. Pass all matplotlib arguments to `curve_kwargs` as a dictionary instead.z>`curve_kwargs` must be None, a dictionary or a list of length z. Got: r:   é   zfTo avoid labeling individual curves that have the same appearance, `curve_kwargs` should be a list of z‹ dictionaries. Alternatively, set `name` to `None` or a single string to label a single legend entry with mean ROC AUC score of all curves.NÚmeanr   Ústdéÿÿÿÿz +/- rB   rC   ÚmetricrG   )r-   ÚwarningsÚwarnÚFutureWarningÚ
isinstanceÚlistr<   Ústrr   Úupdater   rH   ÚextendÚzipÚappendÚ	enumerateÚ_validate_style_kwargs)Ún_curvesr   Úlegend_metricrF   Úcurve_kwargsÚdefault_curve_kwargsÚdefault_multi_curve_kwargsÚkwargsÚlabelsÚlabel_aggregaterD   rE   Úfold_idxrG   Úcurve_kwargs_s                  r   Ú_validate_curve_kwargsz9_BinaryClassifierCurveDisplayMixin._validate_curve_kwargsz   s“  € ñ^ ™FÜð?óð ñ
 Ü�M‰MðRäôð
 "ˆLä�l¤DÔ)¬c°,Ó.?À8Ò.KÜØPØ�*˜G L >°ð4óð ô �tœTÔ"Ü�D“	˜Q’Ü˜|¬TÔ2äð6Ø6>°Zð @OðOóð ô �dœCÔ Ø�6ˆDÜ�dœDÔ!¤c¨$£i°1¢nØ˜(‘?ˆDØ$( L�ˆv˜Ò °dˆô �l¤GÔ,Ø(˜>¨HÑ4‰LØÐ!Ø˜4 (™?ˆLàÐ'Ø#%Ð Ø%Ð-Ø)+Ð&à�aŠ<Ø ×'Ñ'Ð(BÔCàˆØ�]Ñ"Ü@×RÑRØ˜fÑ% t¨A¡wÐ0BóˆOð
 ˜UÑ#Ð/à˜‘7Ð&à'¨¨Ð,°°}ÀUÑ7KÈDÐ6QÐQRÐ/SÑSñ $ð
 (¨E°-ÀÑ2FÀtÐ1LÐ*MÑMð $ð �M‰M˜?Ð+¨t¨f¸À1¹Ñ.EÑEÕFä36°}ÀXÑ7NÐPTÖ3UÑ/Ð# ZØ—‘Ü6×HÑHØ+¨ZÐ9Kóõð 4Vô $-¨VÔ#4ô	
ñ $5‘�˜%ô #Ø˜%Ð8Ð#7Ð8¸,ÀxÑ:Põð $5ð	 	ñ 
ð Ðùó
s   Ç
G,)NN)r   Ú
__module__Ú__qualname__Ú__doc__r   Úclassmethodr'   r/   r@   ÚstaticmethodrH   re   r3   r   r   r   r      s�   „ ñð +/°Tô #ð à17À4Èdó'ó ð'ð  à.2¸dÈóó ðð$ ñ"5ó ð"5ðH ñ
ó ð
ð ð "Ø#'òEó ñEr   r   c                 ó   — | �| S |€|rdS dS t        |«      r|j                  n|} |r| j                  d«      r| dd } nd| › �} n| j                  d«      rd| dd › �} | j                  dd«      } | j	                  «       S )	aÁ  Validate the `score_name` parameter.

    If `score_name` is provided, we just return it as-is.
    If `score_name` is `None`, we use `Score` if `negate_score` is `False` and
    `Negative score` otherwise.
    If `score_name` is a string or a callable, we infer the name. We replace `_` by
    spaces and capitalize the first letter. We remove `neg_` and replace it by
    `"Negative"` if `negate_score` is `False` or just remove it otherwise.
    NzNegative scoreÚScoreÚneg_é   z	Negative r   Ú )Úcallabler   Ú
startswithÚreplaceÚ
capitalize)Ú
score_nameÚscoringÚnegate_scores      r   Ú_validate_score_namerw     s¡   € ð ÐØÐØ	ˆÙ#/ÐÐ<°WÐ<ä)1°'Ô):�W×%Ò%Àˆ
ÙØ×$Ñ$ VÔ,Ø'¨¨˜^‘
à(¨¨Ð5‘
Ø×"Ñ" 6Ô*Ø$ Z°° ^Ð$4Ð5ˆJØ×'Ñ'¨¨SÓ1ˆ
Ø×$Ñ$Ó&Ð&r   c                 ó”   — t        j                  t        j                  | «      «      }|j                  «       |j	                  «       z  S )a   Compute the ratio between the largest and smallest inter-point distances.

    A value larger than 5 typically indicates that the parameter range would
    better be displayed with a log scale while a linear scale would be more
    suitable otherwise.
    )ÚnpÚdiffÚsortÚmaxÚmin)Údatarz   s     r   Ú_interval_max_min_ratior     s1   € ô �7‰7”2—7‘7˜4“=Ó!€DØ�8‰8‹:˜Ÿ™›
Ñ"Ð"r   c                 ób  — i dd“dd“dd“dd“d	d
“dd“dd“dd“dd“dd“dd“dd“dd“dd“dd“dd “d!d"“d#d$d%d&d'd(d)œ¥}|j                  «       D ]   \  }}||v sŒ||v sŒt        d*|› d+|› d,�«      ‚ | j                  «       }|j                  «       D ]  }||v r||   |||   <   Œ||   ||<   Œ |S )-aÃ  Create valid style kwargs by avoiding Matplotlib alias errors.

    Matplotlib raises an error when, for example, 'color' and 'c', or 'linestyle' and
    'ls', are specified together. To avoid this, we automatically keep only the one
    specified by the user and raise an error if the user specifies both.

    Parameters
    ----------
    default_style_kwargs : dict
        The Matplotlib style kwargs used by default in the scikit-learn display.
    user_style_kwargs : dict
        The user-defined Matplotlib style kwargs.

    Returns
    -------
    valid_style_kwargs : dict
        The validated style kwargs taking into account both default and user-defined
        Matplotlib style kwargs.
    ÚlsÚ	linestyleÚcÚcolorÚecÚ	edgecolorÚfcÚ	facecolorÚlwÚ	linewidthÚmecÚmarkeredgecolorÚmfcaltÚmarkerfacecoloraltÚmsÚ
markersizeÚmewÚmarkeredgewidthÚmfcÚmarkerfacecolorÚaaÚantialiasedÚdsÚ	drawstyleÚfontÚfontpropertiesÚfamilyÚ
fontfamilyr   ÚfontnameÚsizeÚfontsizeÚstretchÚfontstretchÚ	fontstyleÚfontvariantÚ
fontweightÚhorizontalalignmentÚverticalalignmentÚmultialignment)ÚstyleÚvariantÚweightÚhaÚvaÚmaz	Got both ú and z", which are aliases of one another)ÚitemsÚ	TypeErrorÚcopyÚkeys)Údefault_style_kwargsÚuser_style_kwargsÚinvalid_to_valid_kwÚinvalid_keyÚ	valid_keyÚvalid_style_kwargsr5   s          r   rZ   rZ   )  sš  € ð*ØˆkðàˆWðð 	ˆkðð 	ˆkð	ð
 	ˆkðð 	Ð ðð 	Ð&ðð 	ˆlðð 	Ð ðð 	Ð ðð 	ˆmðð 	ˆkðð 	Ð ðð 	�,ðð 	�
ðð  	�
ð!ð" 	�=ð#ð$ Ø ØØ#Ø!Øò/Ðð2 #6×";Ñ";Ö"=Ñˆ�YØÐ+Ò+°	Ð=NÒ0NÜØ˜K˜=¨¨i¨[ð 9ð óð ð #>ð .×2Ñ2Ó4Ðà ×%Ñ%Ö'ˆØÐ%Ñ%Ø;LÈSÑ;QÐÐ2°3Ñ7Ò8à&7¸Ñ&<Ð˜sÒ#ð	 (ð Ðr   c                 óš   — dD ]   }| j                   |   j                  d«       Œ" dD ]!  }| j                   |   j                  dd«       Œ# y)z—Remove the top and right spines of the plot.

    Parameters
    ----------
    ax : matplotlib.axes.Axes
        The axes of the plot to despine.
    )ÚtopÚrightF)ÚbottomÚleftr   rJ   N)ÚspinesÚset_visibleÚ
set_bounds)r   Úss     r   Ú_despinerÂ   h  sF   € ó ˆØ
�	‰	�!‰× Ñ  Õ'ð ãˆØ
�	‰	�!‰×Ñ  1Õ%ñ  r   c                 óÐ   — t        |«      }|j                  › d|j                  dz   › �}| dk7  r7|rt        d|› d|› d�«      ‚t	        j
                  d|› d|› d�t        «       | S |S )	z/Deprecate `estimator_name` in favour of `name`.r:   é   Ú
deprecatedzSCannot provide both `estimator_name` and `name`. `estimator_name` is deprecated in ú and will be removed in z. Use `name` only.z"`estimator_name` is deprecated in z. Use `name` instead.)r   ÚmajorÚminorr-   rO   rP   rQ   )r   r   ÚversionÚversion_removes       r   Ú_deprecate_estimator_namerË   v  s•   € ä˜GÓ$€GØŸ™� a¨¯©¸Ñ(9Ð':Ð;€NØ˜Ò%ÙÜð$Ø$+ 9Ð,DÀ^ÐDTð U#ð#óð ô
 	�‰Ø0°°	Ð9QØÐÐ3ð5äô	
ð
 ÐØ€Kr   c                 ó2   — | €yt        | t        «      r| S | gS )z3Convert parameters to a list, leaving `None` as is.N)rR   rS   )Úparams    r   Ú_convert_to_list_leaving_nonerÎ   Š  s    € à€}ØÜ�%œÔØˆØˆ7€Nr   c           	      óü  — i }|j                  «       D ]  \  }}t        |t        «      sŒ|||<   Œ i | ¥|¥}t        |j	                  «       D �ch c]  }t        |«      ’Œ c}«      dkD  r‰|j                  «       D �cg c]  }|‘Œ }}dj                  dj                  |dd «      |d   g«      }	d}
d|v rd}
dj                  d	„ |j                  «       D «       «      }t        |	› d
|› d|
› d|› �«      ‚yc c}w c c}w )z>Check required and optional parameters are of the same length.rJ   r®   z, NrM   Ú z'name' (or self.name)z (or `plot`)c              3   óB   K  — | ]  \  }}|› d t        |«      › �–— Œ y­w)z: N)r<   )r4   r5   Úvalues      r   r7   z'_check_param_lengths.<locals>.<genexpr>£  s(   è ø€ ð &
Ù5G¡z s¨Eˆsˆe�2”c˜%“j�\Ô"Ñ5Gùs   ‚z from `z` initializationz/, should all be lists of the same length. Got: )r¯   rR   rS   r<   Úvaluesr²   Újoinr-   )ÚrequiredÚoptionalÚ
class_nameÚoptional_providedr   rÍ   Ú
all_paramsr5   Ú
param_keysÚparams_formattedÚor_plotÚlengths_formatteds               r   Ú_check_param_lengthsrÞ   “  s,  € àÐØ—~‘~Ö'‰ˆˆeÜ�eœTÕ"Ø&+Ð˜dÒ#ð (ð 3�HÐ2Ð 1Ð2€JÜ
 J×$5Ñ$5Ô$7Ó8Ñ$7˜5ŒC��JÐ$7Ñ8Ó9¸AÒ=Ø%/§_¡_Ô%6Ó7Ñ%6˜c’cÐ%6ˆ
Ð7ð #Ÿ<™<¨¯©°:¸c¸r°?Ó)CÀZÐPRÁ^Ð(TÓUÐØˆØ" jÑ0Ø$ˆGØ ŸI™Iñ &
Ø5?×5EÑ5EÔ5Gó&
ó 
Ðô ØÐ  ¨
 |Ð3CÀGÀ9ð M<Ø<MÐ;NðPó
ð 	
ð >ùÒ8ùÚ7s   ÁC4Á<	C9c                 ó  — t        |«      }|j                  › d|j                  dz   › �}| �'t        |t        «      r|dk(  st        d|› d|› d�«      ‚t        |t        «      r|dk(  s#t        j                  d|› d|› d�t        «       |S | S )z-Deprecate `y_pred` in favour of of `y_score`.r:   rÄ   rÅ   zi`y_pred` and `y_score` cannot be both specified. Please use `y_score` only as `y_pred` was deprecated in rÆ   zy_pred was deprecated in z. Please use `y_score` instead.)	r   rÇ   rÈ   rR   rT   r-   rO   rP   rQ   )Úy_scorer&   rÉ   rÊ   s       r   Ú_deprecate_y_pred_parameterrá   ­  s´   € ä˜GÓ$€GØŸ™� a¨¯©¸Ñ(9Ð':Ð;€NØÐ¤J¨v´sÔ$;ÀÈ,Ò@VÜð3Ø3:°)ð <Ø(Ð)¨ð,ó
ð 	
ô
 �vœsÔ#¨°,Ò(>Ü�‰à+¨G¨9ð 5Ø"Ð#Ð#BðDô ô	
ð ˆà€Nr   )rO   Úcollections.abcr   Únumpyry   Úsklearn.utilsr   Ú$sklearn.utils._optional_dependenciesr   Úsklearn.utils._responser   Úsklearn.utils.fixesr   Úsklearn.utils.multiclassr   Úsklearn.utils.validationr	   r
   r   rw   r   rZ   rÂ   rË   rÎ   rÞ   rá   r3   r   r   Ú<module>rê      sX   ðó Ý #ã å 1Ý IÝ ?Ý -Ý 3ß O÷pñ pòf'ò6#ò<ò~&òò(ò
ó4r   