o
    Ù­jÝR  ã                   @   s\  d 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mZmZmZmZmZ ddlZddlmZ ddlmZmZ ddlmZmZmZmZmZ g d	¢Zeeeeef f Z eZ!G d
d„ deƒZ"de
e# de
ee#eef  fdd„Z$edƒZ%de%de%fdd„Z&G dd„ dƒZ'G dd„ de"ƒZ(G dd„ de"ƒZ)G dd„ de"ƒZ*G dd„ de"ƒZ+dS )z}Callback library containing training routines.  See :doc:`Callback Functions
</python/callbacks>` for a quick introduction.

é    N)ÚABC)ÚAnyÚCallableÚDictÚListÚOptionalÚSequenceÚTupleÚ	TypeAliasÚTypeVarÚUnionÚcasté   )Ú
collective)ÚEvalsLogÚ
_ScoreList)ÚBoosterÚDMatrixÚXGBoostErrorÚ_deprecate_positional_argsÚ_parse_eval_str)ÚTrainingCallbackÚLearningRateSchedulerÚEarlyStoppingÚEvaluationMonitorÚTrainingCheckPointÚCallbackContainerc                   @   s€   e Zd ZU dZeZeed< ddd„Zdedefdd	„Z	dedefd
d„Z
dedededefdd„Zdedededefdd„ZdS )r   zCInterface for training callback.

    .. versionadded:: 1.3.0

    r   ÚreturnNc                 C   s   d S ©N© )Úselfr   r   úM/var/www/html/CropPilot/venv/lib/python3.10/site-packages/xgboost/callback.pyÚ__init__<   s   zTrainingCallback.__init__Úmodelc                 C   ó   |S )zRun before training starts.r   ©r    r#   r   r   r!   Úbefore_training?   ó   z TrainingCallback.before_trainingc                 C   r$   )zRun after training is finished.r   r%   r   r   r!   Úafter_trainingC   r'   zTrainingCallback.after_trainingÚepochÚ	evals_logc                 C   ó   dS )z�Run before each iteration.  Returns True when training should stop. See
        :py:meth:`after_iteration` for details.

        Fr   ©r    r#   r)   r*   r   r   r!   Úbefore_iterationG   s   z!TrainingCallback.before_iterationc                 C   r+   )añ  Run after each iteration.  Returns `True` when training should stop.

        Parameters
        ----------

        model :
            Eeither a :py:class:`~xgboost.Booster` object or a CVPack if the cv function
            in xgboost is being used.
        epoch :
            The current training iteration.
        evals_log :
            A dictionary containing the evaluation history:

            .. code-block:: python

                {"data_name": {"metric_name": [0.5, ...]}}

        Fr   r,   r   r   r!   Úafter_iterationN   s   z TrainingCallback.after_iteration)r   N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r
   Ú__annotations__r"   Ú_Modelr&   r(   ÚintÚboolr-   r.   r   r   r   r!   r   3   s   
 
r   Úrlistr   c                 C   s  i }| d   ¡ d }| D ]B}|  ¡ }||d ksJ ‚t|dd… ƒD ]+\}}t|tƒs/| ¡ }|  d¡\}}||f|vrBg |||f< |||f  t|ƒ¡ q"q|}	g }
t| ¡ dd„ d�D ](\\}}}t	 
|¡}t|	tƒsq|	 ¡ }	t	 |¡t	 |¡}}|
 |||fg¡ q]|
S )z#Aggregate cross-validation results.r   r   Nú:c                 S   s   | d d S ©Nr   r   )Úxr   r   r!   Ú<lambda>u   s    z_aggcv.<locals>.<lambda>)Úkey)ÚsplitÚ	enumerateÚ
isinstanceÚstrÚdecodeÚappendÚfloatÚsortedÚitemsÚnumpyÚarrayÚmeanÚstdÚextend)r7   ÚcvmapÚidxÚlineÚarrÚ
metric_idxÚitÚkÚvÚmsgÚresultsÚ_ÚnameÚsÚas_arrrH   rI   r   r   r!   Ú_aggcvd   s,   
ú 

rY   Ú_ARTÚscorec                 C   sZ   t  ¡ }|dks
J ‚|dkr| S t| tƒrtdƒ‚t | g¡}t  |t jj	¡| }|d S )z§Helper function for computing customized metric in distributed
    environment.  Not strictly correct as many functions don't use mean value
    as final result.

    r   r   zBxgboost.cv function should not be used in distributed environment.)
r   Úget_world_sizer?   ÚtupleÚ
ValueErrorrF   rG   Ú	allreduceÚOpÚSUM)r[   ÚworldrN   r   r   r!   Ú_allreduce_metric‚   s   
ÿrc   c                   @   sö   e Zd ZdZ			ddee dee deded	df
d
d„Z	de
d	e
fdd„Zde
d	e
fdd„Zde
dededeeeeef   d	ef
dd„Zdeeeeef  eeeeef  f ded	dfdd„Zde
dededeeeeef   d	ef
dd„ZdS )r   zfA special internal callback for invoking a list of other callbacks.

    .. versionadded:: 1.3.0

    NTFÚ	callbacksÚmetricÚoutput_marginÚis_cvr   c                 C   sx   t t |¡ƒ| _|D ]}t|tƒstdƒ‚q
d}|d ur$t|ƒs$t|ƒ‚|| _t	 
¡ | _|| _|| _| jr:d | _d S d S )Nz3callback must be an instance of `TrainingCallback`.z†metric must be callable object for monitoring.  For builtin metrics, passing them in training parameter invokes monitor automatically.)ÚlistÚdictÚfromkeysrd   r?   r   Ú	TypeErrorÚcallablere   ÚcollectionsÚOrderedDictÚhistoryÚ_output_marginrg   Úaggregated_cv)r    rd   re   rf   rg   ÚcbrS   r   r   r!   r"   œ   s    
ÿÿ

ÿzCallbackContainer.__init__r#   c                 C   óN   | j D ]!}|j|d�}d}| jrt|jtƒsJ |ƒ‚qt|tƒs$J |ƒ‚q|S )z Function called before training.©r#   z'before_training should return the model)rd   r&   rg   r?   Úcvfoldsrh   r   ©r    r#   ÚcrS   r   r   r!   r&   ·   s   
z!CallbackContainer.before_trainingc                 C   rs   )zFunction called after training.rt   z&after_training should return the model)rd   r(   rg   r?   ru   rh   r   rv   r   r   r!   r(   Â   s   
z CallbackContainer.after_trainingr)   ÚdtrainÚevalsc                    s   t ‡ ‡‡fdd„ˆjD ƒƒS )z*Function called before training iteration.c                 3   ó    � | ]}|  ˆˆ ˆj¡V  qd S r   )r-   ro   ©Ú.0rw   ©r)   r#   r    r   r!   Ú	<genexpr>Ö   s   € 
ÿz5CallbackContainer.before_iteration.<locals>.<genexpr>)Úanyrd   )r    r#   r)   rx   ry   r   r}   r!   r-   Î   s   ÿz"CallbackContainer.before_iterationr[   c                 C   s  |D ]~}|d }|d }| j r"ttttttf |ƒd ƒ}||f}n|}| d¡}|d }	d |dd … ¡}
t|ƒ}|	| jvrFt	 
¡ | j|	< | j|	 }|
|vrVttg ƒ||
< ||
 }| j rstttttf  |ƒ ttttf |ƒ¡ qttt |ƒ tt|ƒ¡ qd S )Nr   r   é   ú-)rg   rC   r   r	   r@   r=   Újoinrc   ro   rm   rn   r   r   rB   )r    r[   r)   ÚdrV   rW   rI   r:   Úsplited_namesÚ	data_nameÚmetric_nameÚdata_historyÚmetric_historyr   r   r!   Ú_update_historyÚ   s.   



ÿéz!CallbackContainer._update_historyc                    s°   ˆj rˆ ˆ ˆjˆj¡}t|ƒ}|ˆ_ˆ |ˆ ¡ n.|du r g n|}|D ]\}}| d¡dks3J dƒ‚q$ˆ |ˆ ˆjˆj¡}t	|ƒ}	ˆ |	ˆ ¡ t
‡ ‡‡fdd„ˆjD ƒƒ}
|
S )z)Function called after training iteration.Nr�   éÿÿÿÿz#Dataset name should not contain `-`c                 3   rz   r   )r.   ro   r{   r}   r   r!   r~     s   € z4CallbackContainer.after_iteration.<locals>.<genexpr>)rg   Úevalre   rp   rY   rq   r‰   ÚfindÚeval_setr   r   rd   )r    r#   r)   rx   ry   ÚscoresrU   rV   r[   Úmetric_scoreÚretr   r}   r!   r.   ø   s   z!CallbackContainer.after_iteration)NTF)r/   r0   r1   r2   r   r   r   r   r6   r"   r4   r&   r(   r5   r   r   r	   r@   r-   r   rC   r‰   r.   r   r   r   r!   r   •   s^    	ûþýüû
úþýüû
ú$þý
üþýüûúr   c                       sZ   e Zd ZdZdeeegef ee f ddf‡ fdd„Z	de
ded	edefd
d„Z‡  ZS )r   a  Callback function for scheduling learning rate.

    .. versionadded:: 1.3.0

    Parameters
    ----------

    learning_rates :
        If it's a callable object, then it should accept an integer parameter
        `epoch` and returns the corresponding learning rate.  Otherwise it
        should be a sequence like list or tuple with the same size of boosting
        rounds.

    Úlearning_ratesr   Nc                    sT   t ˆ ƒstˆ tjjƒstdtˆ ƒ› �ƒ‚t ˆ ƒrˆ | _n‡ fdd„| _tƒ  	¡  d S )Nz=Invalid learning rates, expecting callable or sequence, got: c                    s   t tˆ ƒ|  S r   )r   r   )r)   ©r‘   r   r!   r;   .  s    z0LearningRateScheduler.__init__.<locals>.<lambda>)
rl   r?   rm   Úabcr   rk   Útyper‘   Úsuperr"   )r    r‘   ©Ú	__class__r’   r!   r"      s   
ÿÿÿzLearningRateScheduler.__init__r#   r)   r*   c                 C   s   |  d|  |¡¡ dS )NÚlearning_rateF)Ú	set_paramr‘   r,   r   r   r!   r.   1  s   z%LearningRateScheduler.after_iteration)r/   r0   r1   r2   r   r   r5   rC   r   r"   r4   r   r6   r.   Ú__classcell__r   r   r–   r!   r     s    ÿþ"r   c                       sÀ   e Zd ZdZeddddddœdedee dee d	ee d
ee de	ddf‡ fdd„ƒZ
dedefdd„Zdedededededefdd„Zdedededefdd„Zdedefdd„Z‡  ZS )r   aU  Callback function for early stopping

    .. versionadded:: 1.3.0

    Parameters
    ----------
    rounds :
        Early stopping rounds.
    metric_name :
        Name of metric that is used for early stopping.
    data_name :
        Name of dataset that is used for early stopping.
    maximize :
        Whether to maximize evaluation metric.  None means auto (discouraged).
    save_best :
        Whether training should return the best model or the last model. If set to
        `True`, it will only keep the boosting rounds up to the detected best iteration,
        discarding the ones that come after. This is only supported with tree methods
        (not `gblinear`). Also, the `cv` function doesn't return a model, the parameter
        is not applicable.
    min_delta :

        .. versionadded:: 1.5.0

        Minimum absolute change in score to be qualified as an improvement.

    Examples
    --------

    .. code-block:: python

        es = xgboost.callback.EarlyStopping(
            rounds=2,
            min_delta=1e-3,
            save_best=True,
            maximize=False,
            data_name="validation_0",
            metric_name="mlogloss",
        )
        clf = xgboost.XGBClassifier(tree_method="hist", device="cuda", callbacks=[es])

        X, y = load_digits(return_X_y=True)
        clf.fit(X, y, eval_set=[(X, y)])
    NFg        )r†   r…   ÚmaximizeÚ	save_bestÚ	min_deltaÚroundsr†   r…   r›   rœ   r�   r   c                   s\   || _ || _|| _|| _|| _i | _|| _| jdk rtdƒ‚d| _i | _	d| _
tƒ  ¡  d S )Nr   z(min_delta must be greater or equal to 0.)Údatar†   rž   rœ   r›   Ústopping_historyÚ
_min_deltar^   Úcurrent_roundsÚbest_scoresÚstarting_roundr•   r"   )r    rž   r†   r…   r›   rœ   r�   r–   r   r!   r"   f  s   
zEarlyStopping.__init__r#   c                 C   s&   |  ¡ | _t|tƒs| jrtdƒ‚|S )NzP`save_best` is not applicable to the `cv` function as it doesn't return a model.)Únum_boosted_roundsr¤   r?   r   rœ   r^   r%   r   r   r!   r&   €  s   
ÿzEarlyStopping.before_trainingr[   rV   re   r)   c                   s   dt dtfdd„‰ dt dt dtf‡ ‡fdd„}dt dt dtf‡ ‡fd	d
„}ˆjd u rBd}ˆdkr?t‡fdd„|D ƒƒr?dˆ_ndˆ_ˆjrH|}	n|}	ˆjs{dˆ_i ˆj|< tt|gƒˆj| ˆ< i ˆj	|< |gˆj	| ˆ< |j
tˆ |ƒƒt|ƒd� nK|	|ˆj	| ˆ d ƒs™ˆj| ˆ  |¡ ˆ jd7  _n-ˆj| ˆ  |¡ ˆj	| ˆ  |¡ ˆj| ˆ d }
|j
tˆ |
ƒƒt|ƒd� dˆ_ˆjˆjkrÎdS dS )NÚvaluer   c                 S   s   t | tƒr	| d S | S )z+get score if it's cross validation history.r   )r?   r]   )r¦   r   r   r!   Úget_sŒ  s   z+EarlyStopping._update_rounds.<locals>.get_sÚnewÚbestc                    s   t  ˆ | ƒˆj ˆ |ƒ¡S )z-New score should be greater than the old one.©rF   Úgreaterr¡   ©r¨   r©   ©r§   r    r   r!   r›   �  ó   z.EarlyStopping._update_rounds.<locals>.maximizec                    s   t  ˆ |ƒˆj ˆ | ƒ¡S )z,New score should be lesser than the old one.rª   r¬   r­   r   r!   Úminimize”  r®   z.EarlyStopping._update_rounds.<locals>.minimize)
ÚaucÚaucprÚprezpre@ÚmapÚndcgzauc@zaucpr@zmap@zndcg@Úmapec                 3   s   � | ]}ˆ   |¡V  qd S r   )Ú
startswith)r|   r:   )re   r   r!   r~   §  s   € z/EarlyStopping._update_rounds.<locals>.<genexpr>TFr   )Ú
best_scoreÚbest_iterationrŠ   r   )Ú_ScorerC   r6   r›   r   r    r¢   r   r   r£   Úset_attrr@   rB   rž   )r    r[   rV   re   r#   r)   r›   r¯   Úmaximize_metricsÚ
improve_opÚrecordr   )r§   re   r    r!   Ú_update_rounds‰  s:   


zEarlyStopping._update_roundsr*   c           	      C   sÒ   || j 7 }d}t| ¡ ƒdk rt|ƒ‚| jr| j}nt| ¡ ƒd }||vr-td|› �ƒ‚t|tƒs;tdt	|ƒ› �ƒ‚|| }| j
rF| j
}nt| ¡ ƒd }||vrYtd|› �ƒ‚|| d }| j|||||d�S )Nz;Must have at least 1 validation dataset for early stopping.r   rŠ   zNo dataset named: z1The name of the dataset should be a string. Got: zNo metric named: )r[   rV   re   r#   r)   )r¤   ÚlenÚkeysr^   rŸ   rh   r?   r@   rk   r”   r†   r¾   )	r    r#   r)   r*   rS   r…   Údata_logr†   r[   r   r   r!   r.   È  s.   

ÿ
ÿzEarlyStopping.after_iterationc              
   C   sp   | j s|S z!|j}|j}|d ur|d usJ ‚|d |d … }||_||_W |S  ty7 } ztdƒ|‚d }~ww )Nr   z4`save_best` is not applicable to the current booster)rœ   r¸   r·   r   )r    r#   r¸   r·   Úer   r   r!   r(   ì  s$   ûÿþ€ÿzEarlyStopping.after_training)r/   r0   r1   r2   r   r5   r   r@   r6   rC   r"   r4   r&   r¹   r¾   r   r.   r(   rš   r   r   r–   r!   r   7  sN    .øýüûúùø	÷	ÿÿÿÿÿ
þ?$r   c                       s–   e Zd ZdZdddejfdedededee	gd	f f‡ fd
d„Z
de	de	dedee de	f
dd„Zdedededefdd„Zdedefdd„Z‡  ZS )r   a’  Print the evaluation result at each iteration.

    .. versionadded:: 1.3.0

    Parameters
    ----------

    rank :
        Which worker should be used for printing the result.
    period :
        How many epoches between printing.
    show_stdv :
        Used in cv to show standard deviation.  Users should not specify it.
    logger :
        A callable used for logging evaluation result.

    r   r   FÚrankÚperiodÚ	show_stdvÚloggerNc                    s8   || _ || _|| _|| _|dksJ ‚d | _tƒ  ¡  d S r9   )Úprinter_rankrÅ   rÄ   Ú_loggerÚ_latestr•   r"   )r    rÃ   rÄ   rÅ   rÆ   r–   r   r!   r"     s   zEvaluationMonitor.__init__rŸ   re   r[   rI   r   c                 C   sR   |d ur| j rd|d | › d|d›d|d›�}|S d|d | › d|d›�}|S )Nú	r�   r8   z.5fú+)rÅ   )r    rŸ   re   r[   rI   rS   r   r   r!   Ú_fmt_metric"  s
   "ÿzEvaluationMonitor._fmt_metricr#   r)   r*   c              	   C   sÌ   |sdS d|› d�}t  ¡ | jkrd| ¡ D ]1\}}| ¡ D ](\}}d }	t|d tƒr7|d d }
|d d }	n|d }
||  |||
|	¡7 }qq|d7 }|| j dksW| jdkra|  |¡ d | _	dS || _	dS )NFú[ú]rŠ   r   r   Ú
)
r   Úget_rankrÇ   rE   r?   r]   rÌ   rÄ   rÈ   rÉ   )r    r#   r)   r*   rS   rŸ   re   r†   ÚlogÚstdvr[   r   r   r!   r.   +  s(   ù
ÿz!EvaluationMonitor.after_iterationc                 C   s(   t  ¡ | jkr| jd ur|  | j¡ |S r   )r   rÐ   rÇ   rÉ   rÈ   r%   r   r   r!   r(   D  s   z EvaluationMonitor.after_training)r/   r0   r1   r2   r   Úcommunicator_printr5   r6   r   r@   r"   rC   r   rÌ   r4   r   r.   r(   rš   r   r   r–   r!   r   ÿ  s8    ûþýüûÿÿÿÿ
þ	r   c                       sx   e Zd ZdZdZ			ddeeejf dede	d	e
d
df
‡ fdd„Zded
efdd„Zdede
ded
e	fdd„Z‡  ZS )r   ak  Checkpointing operation. Users are encouraged to create their own callbacks for
    checkpoint as XGBoost doesn't handle distributed file systems. When checkpointing on
    distributed systems, be sure to know the rank of the worker to avoid multiple
    workers checkpointing to the same place.

    .. versionadded:: 1.3.0

    Since XGBoost 2.1.0, the default format is changed to UBJSON.

    Parameters
    ----------

    directory :
        Output model directory.
    name :
        pattern of output model file.  Models will be saved as name_0.ubj, name_1.ubj,
        name_2.ubj ....
    as_pickle :
        When set to True, all training parameters will be saved in pickle format,
        instead of saving only the model.
    interval :
        Interval of checkpointing.  Checkpointing is slow so setting a larger number can
        reduce performance hit.

    Úubjr#   Féd   Ú	directoryrV   Ú	as_pickleÚintervalr   Nc                    s8   t  |¡| _|| _|| _|| _d| _d| _tƒ  	¡  d S r9   )
ÚosÚfspathÚ_pathÚ_nameÚ
_as_pickleÚ_iterationsÚ_epochÚ_startr•   r"   )r    rÖ   rV   r×   rØ   r–   r   r!   r"   g  s   zTrainingCheckPoint.__init__c                 C   s   |  ¡ | _|S r   )r¥   rà   r%   r   r   r!   r&   v  s   
z"TrainingCheckPoint.before_trainingr)   r*   c                 C   s²   | j | jkrPtj | j| jd t|| j ƒ | j	rdnd| j
› � ¡}d| _ t ¡ dkrP| j	rKt|dƒ�}t ||¡ W d   ƒ n1 sEw   Y  n| |¡ |  j d7  _ dS )NrU   z.pklÚ.r   Úwbr   F)rß   rÞ   rÙ   Úpathr‚   rÛ   rÜ   r@   rà   rÝ   Údefault_formatr   rÐ   ÚopenÚpickleÚdumpÚ
save_model)r    r#   r)   r*   rã   Úfdr   r   r!   r.   z  s*   ÿþýþÿ€
z"TrainingCheckPoint.after_iteration)r#   FrÕ   )r/   r0   r1   r2   rä   r   r@   rÙ   ÚPathLiker6   r5   r"   r4   r&   r   r.   rš   r   r   r–   r!   r   J  s&    ûþýüûú"r   ),r2   rm   rÙ   ræ   r“   r   Útypingr   r   r   r   r   r   r	   r
   r   r   r   rF   Ú r   Ú_typingr   r   Úcorer   r   r   r   r   Ú__all__rC   r¹   r4   r   r@   rY   rZ   rc   r   r   r   r   r   r   r   r   r!   Ú<module>   s.    4	$1{' IK