U
    dix                     @   s  d 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	Z
ddlmZ ddlmZmZ ddlmZ dd	lmZ eeZeG d
d dZG dd dZeG dd de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G dd deeZdS )zJ
Callbacks to use with the Trainer class and customize the training loop.
    N)	dataclass)DictListOptionalUnion)tqdm   )IntervalStrategy
has_length)TrainingArguments)loggingc                   @   sN  e Zd ZU dZdZee ed< dZe	ed< dZ
e	ed< dZe	ed< dZe	ed	< dZe	ed
< dZe	ed< dZe	ed< dZe	ed< dZeed< dZeeeef  ed< dZee ed< dZee ed< dZeed< dZeed< dZeed< dZeed< dZeeeeee	ef f ed< dZed ed< dd Z edddZ!e"edd d!Z#dS )"TrainerStatea  
    A class containing the [`Trainer`] inner state that will be saved along the model and optimizer when checkpointing
    and passed to the [`TrainerCallback`].

    <Tip>

    In all this class, one step is to be understood as one update step. When using gradient accumulation, one update
    step may require several forward and backward passes: if you use `gradient_accumulation_steps=n`, then one update
    step requires going through *n* batches.

    </Tip>

    Args:
        epoch (`float`, *optional*):
            Only set during training, will represent the epoch the training is at (the decimal part being the
            percentage of the current epoch completed).
        global_step (`int`, *optional*, defaults to 0):
            During training, represents the number of update steps completed.
        max_steps (`int`, *optional*, defaults to 0):
            The number of update steps to do during the current training.
        logging_steps (`int`, *optional*, defaults to 500):
            Log every X updates steps
        eval_steps (`int`, *optional*):
            Run an evaluation every X steps.
        save_steps (`int`, *optional*, defaults to 500):
            Save checkpoint every X updates steps.
        train_batch_size (`int`, *optional*):
            The batch size for the training dataloader. Only needed when
            `auto_find_batch_size` has been used.
        num_input_tokens_seen (`int`, *optional*, defaults to 0):
            The number of tokens seen during training (number of input tokens, not the number of prediction tokens).
        total_flos (`float`, *optional*, defaults to 0):
            The total number of floating operations done by the model since the beginning of training (stored as floats
            to avoid overflow).
        log_history (`List[Dict[str, float]]`, *optional*):
            The list of logs done since the beginning of training.
        best_metric (`float`, *optional*):
            When tracking the best model, the value of the best metric encountered so far.
        best_model_checkpoint (`str`, *optional*):
            When tracking the best model, the value of the name of the checkpoint for the best model encountered so
            far.
        is_local_process_zero (`bool`, *optional*, defaults to `True`):
            Whether or not this process is the local (e.g., on one machine if training in a distributed fashion on
            several machines) main process.
        is_world_process_zero (`bool`, *optional*, defaults to `True`):
            Whether or not this process is the global main process (when training in a distributed fashion on several
            machines, this is only going to be `True` for one process).
        is_hyper_param_search (`bool`, *optional*, defaults to `False`):
            Whether we are in the process of a hyper parameter search using Trainer.hyperparameter_search. This will
            impact the way data will be logged in TensorBoard.
        stateful_callbacks (`List[StatefulTrainerCallback]`, *optional*):
            Callbacks attached to the `Trainer` that should have their states be saved or restored.
            Relevent callbacks should implement a `state` and `from_state` function.
    Nepochr   global_step	max_stepsi  logging_steps
eval_steps
save_stepstrain_batch_sizenum_train_epochsnum_input_tokens_seen
total_floslog_historybest_metricbest_model_checkpointTis_local_process_zerois_world_process_zeroFis_hyper_param_search
trial_nametrial_paramsTrainerCallbackstateful_callbacksc                 C   s   | j d krg | _ | jd kr"i | _nt| jtr0n~i }| jD ]l}t|tsZtdt| |jj}||krt|| t	s|| g||< || 
|  q:| ||< q:|| _d S )NzNAll callbacks passed to be saved must inherit `ExportableState`, but received )r   r!   
isinstancedictExportableState	TypeErrortype	__class____name__listappendstate)selfr!   callbackname r/   A/tmp/pip-unpacked-wheel-bm_b0l5e/transformers/trainer_callback.py__post_init__p   s&    



zTrainerState.__post_init__)	json_pathc              	   C   sB   t jt| dddd }t|ddd}|| W 5 Q R X dS )	zDSave the content of this instance in JSON format inside `json_path`.   T)indent	sort_keys
wutf-8encodingN)jsondumpsdataclassesZasdictopenwrite)r,   r2   Zjson_stringfr/   r/   r0   save_to_json   s    zTrainerState.save_to_jsonc              	   C   s2   t |ddd}| }W 5 Q R X | f t|S )z3Create an instance from the content of `json_path`.rr8   r9   )r>   readr;   loads)clsr2   r@   textr/   r/   r0   load_from_json   s    zTrainerState.load_from_json)$r(   
__module____qualname____doc__r   r   float__annotations__r   intr   r   r   r   r   r   r   r   r   r   r   strr   r   r   boolr   r   r   r   r   r!   r1   rA   classmethodrG   r/   r/   r/   r0   r   #   s0   
7 r   c                   @   s*   e Zd ZdZedddZedd ZdS )r$   aj  
    A class for objects that include the ability to have its state
    be saved during `Trainer._save_checkpoint` and loaded back in during
    `Trainer._load_from_checkpoint`.

    These must implement a `state` function that gets called during the respective
    Trainer function call. It should only include parameters and attributes needed to
    recreate the state at a particular time, to avoid utilizing pickle/maintain standard
    file IO writing.

    Example:

    ```python
    class EarlyStoppingCallback(TrainerCallback, ExportableState):
        def __init__(self, early_stopping_patience: int = 1, early_stopping_threshold: Optional[float] = 0.0):
            self.early_stopping_patience = early_stopping_patience
            self.early_stopping_threshold = early_stopping_threshold
            # early_stopping_patience_counter denotes the number of times validation metrics failed to improve.
            self.early_stopping_patience_counter = 0

        def state(self) -> dict:
            return {
                "args": {
                    "early_stopping_patience": self.early_stopping_patience,
                    "early_stopping_threshold": self.early_stopping_threshold,
                },
                "attributes": {
                    "early_stopping_patience_counter": self.early_stopping_patience_counter,
                }
            }
    ```returnc                 C   s   t dd S )Nz<You must implement a `state` function to utilize this class.)NotImplementedErrorr,   r/   r/   r0   r+      s    zExportableState.statec                 C   s4   | f |d }|d   D ]\}}t||| q|S )Nargs
attributes)itemssetattr)rE   r+   instancekvr/   r/   r0   
from_state   s    zExportableState.from_stateN)r(   rH   rI   rJ   r#   r+   rP   r\   r/   r/   r/   r0   r$      s    r$   c                   @   st   e Zd ZU dZdZeed< dZeed< dZeed< dZ	eed< dZ
eed< dd	 Zd
d Zdd ZedddZdS )TrainerControlaA  
    A class that handles the [`Trainer`] control flow. This class is used by the [`TrainerCallback`] to activate some
    switches in the training loop.

    Args:
        should_training_stop (`bool`, *optional*, defaults to `False`):
            Whether or not the training should be interrupted.

            If `True`, this variable will not be set back to `False`. The training will just stop.
        should_epoch_stop (`bool`, *optional*, defaults to `False`):
            Whether or not the current epoch should be interrupted.

            If `True`, this variable will be set back to `False` at the beginning of the next epoch.
        should_save (`bool`, *optional*, defaults to `False`):
            Whether or not the model should be saved at this step.

            If `True`, this variable will be set back to `False` at the beginning of the next step.
        should_evaluate (`bool`, *optional*, defaults to `False`):
            Whether or not the model should be evaluated at this step.

            If `True`, this variable will be set back to `False` at the beginning of the next step.
        should_log (`bool`, *optional*, defaults to `False`):
            Whether or not the logs should be reported at this step.

            If `True`, this variable will be set back to `False` at the beginning of the next step.
    Fshould_training_stopshould_epoch_stopshould_saveshould_evaluate
should_logc                 C   s
   d| _ dS )z<Internal method that resets the variable for a new training.FN)r^   rT   r/   r/   r0   _new_training   s    zTrainerControl._new_trainingc                 C   s
   d| _ dS )z9Internal method that resets the variable for a new epoch.FN)r_   rT   r/   r/   r0   
_new_epoch   s    zTrainerControl._new_epochc                 C   s   d| _ d| _d| _dS )z8Internal method that resets the variable for a new step.FN)r`   ra   rb   rT   r/   r/   r0   	_new_step   s    zTrainerControl._new_steprQ   c                 C   s    | j | j| j| j| jdi dS )Nr^   r_   r`   ra   rb   rU   rV   rf   rT   r/   r/   r0   r+      s    zTrainerControl.stateN)r(   rH   rI   rJ   r^   rO   rL   r_   r`   ra   rb   rc   rd   re   r#   r+   r/   r/   r/   r0   r]      s   
r]   c                   @   s  e Zd ZdZeeedddZeeedddZeeedddZ	eeedd	d
Z
eeedddZeeedddZeeedddZeeedddZeeedddZeeedddZeeedddZeeedddZeeedddZeeedddZeeeddd Zd!S )"r    a	  
    A class for objects that will inspect the state of the training loop at some events and take some decisions. At
    each of those events the following arguments are available:

    Args:
        args ([`TrainingArguments`]):
            The training arguments used to instantiate the [`Trainer`].
        state ([`TrainerState`]):
            The current state of the [`Trainer`].
        control ([`TrainerControl`]):
            The object that is returned to the [`Trainer`] and can be used to make some decisions.
        model ([`PreTrainedModel`] or `torch.nn.Module`):
            The model being trained.
        tokenizer ([`PreTrainedTokenizer`]):
            The tokenizer used for encoding the data. This is deprecated in favour of `processing_class`.
        processing_class ([`PreTrainedTokenizer` or `BaseImageProcessor` or `ProcessorMixin` or `FeatureExtractionMixin`]):
            The processing class used for encoding the data. Can be a tokenizer, a processor, an image processor or a feature extractor.
        optimizer (`torch.optim.Optimizer`):
            The optimizer used for the training steps.
        lr_scheduler (`torch.optim.lr_scheduler.LambdaLR`):
            The scheduler used for setting the learning rate.
        train_dataloader (`torch.utils.data.DataLoader`, *optional*):
            The current dataloader used for training.
        eval_dataloader (`torch.utils.data.DataLoader`, *optional*):
            The current dataloader used for evaluation.
        metrics (`Dict[str, float]`):
            The metrics computed by the last evaluation phase.

            Those are only accessible in the event `on_evaluate`.
        logs  (`Dict[str, float]`):
            The values to log.

            Those are only accessible in the event `on_log`.

    The `control` object is the only one that can be changed by the callback, in which case the event that changes it
    should return the modified version.

    The argument `args`, `state` and `control` are positionals for all events, all the others are grouped in `kwargs`.
    You can unpack the ones you need in the signature of the event using them. As an example, see the code of the
    simple [`~transformers.PrinterCallback`].

    Example:

    ```python
    class PrinterCallback(TrainerCallback):
        def on_log(self, args, state, control, logs=None, **kwargs):
            _ = logs.pop("total_flos", None)
            if state.is_local_process_zero:
                print(logs)
    ```rU   r+   controlc                 K   s   dS )zS
        Event called at the end of the initialization of the [`Trainer`].
        Nr/   r,   rU   r+   ri   kwargsr/   r/   r0   on_init_end8  s    zTrainerCallback.on_init_endc                 K   s   dS )z<
        Event called at the beginning of training.
        Nr/   rj   r/   r/   r0   on_train_begin>  s    zTrainerCallback.on_train_beginc                 K   s   dS )z6
        Event called at the end of training.
        Nr/   rj   r/   r/   r0   on_train_endD  s    zTrainerCallback.on_train_endc                 K   s   dS )z<
        Event called at the beginning of an epoch.
        Nr/   rj   r/   r/   r0   on_epoch_beginJ  s    zTrainerCallback.on_epoch_beginc                 K   s   dS )z6
        Event called at the end of an epoch.
        Nr/   rj   r/   r/   r0   on_epoch_endP  s    zTrainerCallback.on_epoch_endc                 K   s   dS )z
        Event called at the beginning of a training step. If using gradient accumulation, one training step might take
        several inputs.
        Nr/   rj   r/   r/   r0   on_step_beginV  s    zTrainerCallback.on_step_beginc                 K   s   dS )zv
        Event called before the optimizer step but after gradient clipping. Useful for monitoring gradients.
        Nr/   rj   r/   r/   r0   on_pre_optimizer_step]  s    z%TrainerCallback.on_pre_optimizer_stepc                 K   s   dS )z}
        Event called after the optimizer step but before gradients are zeroed out. Useful for monitoring gradients.
        Nr/   rj   r/   r/   r0   on_optimizer_stepc  s    z!TrainerCallback.on_optimizer_stepc                 K   s   dS )zU
        Event called at the end of an substep during gradient accumulation.
        Nr/   rj   r/   r/   r0   on_substep_endi  s    zTrainerCallback.on_substep_endc                 K   s   dS )z
        Event called at the end of a training step. If using gradient accumulation, one training step might take
        several inputs.
        Nr/   rj   r/   r/   r0   on_step_endo  s    zTrainerCallback.on_step_endc                 K   s   dS )z9
        Event called after an evaluation phase.
        Nr/   rj   r/   r/   r0   on_evaluatev  s    zTrainerCallback.on_evaluatec                 K   s   dS )z=
        Event called after a successful prediction.
        Nr/   )r,   rU   r+   ri   metricsrk   r/   r/   r0   
on_predict|  s    zTrainerCallback.on_predictc                 K   s   dS )z7
        Event called after a checkpoint save.
        Nr/   rj   r/   r/   r0   on_save  s    zTrainerCallback.on_savec                 K   s   dS )z;
        Event called after logging the last logs.
        Nr/   rj   r/   r/   r0   on_log  s    zTrainerCallback.on_logc                 K   s   dS )z7
        Event called after a prediction step.
        Nr/   rj   r/   r/   r0   on_prediction_step  s    z"TrainerCallback.on_prediction_stepN)r(   rH   rI   rJ   r   r   r]   rl   rm   rn   ro   rp   rq   rr   rs   rt   ru   rv   rx   ry   rz   r{   r/   r/   r/   r0   r      s    3r    c                   @   sR  e Zd ZdZdd Zdd Zdd Zdd	 Zed
d Z	e
eedddZe
eedddZe
eedddZe
eedddZe
eedddZe
eedddZe
eedddZe
eedddZe
eedddZe
eeddd Ze
eedd!d"Ze
eedd#d$Ze
eedd%d&Ze
eedd'd(Ze
eedd)d*Zd+d, Zd-S ).CallbackHandlerz>Internal class that just calls the list of callbacks in order.c                 C   sf   g | _ |D ]}| | q
|| _|| _|| _|| _d | _d | _tdd | j D sbt	
d| j  d S )Nc                 s   s   | ]}t |tV  qd S N)r"   DefaultFlowCallback.0cbr/   r/   r0   	<genexpr>  s     z+CallbackHandler.__init__.<locals>.<genexpr>zThe Trainer will not work properly if you don't have a `DefaultFlowCallback` in its callbacks. You
should add one before training with `trainer.add_callback(DefaultFlowCallback). The current list ofcallbacks is
:)	callbacksadd_callbackmodelprocessing_class	optimizerlr_schedulertrain_dataloadereval_dataloaderanyloggerwarningcallback_list)r,   r   r   r   r   r   r   r/   r/   r0   __init__  s    zCallbackHandler.__init__c                 C   sh   t |tr| n|}t |tr"|n|j}|dd | jD krXtd| dd | j  | j| d S )Nc                 S   s   g | ]
}|j qS r/   )r'   )r   cr/   r/   r0   
<listcomp>  s     z0CallbackHandler.add_callback.<locals>.<listcomp>zYou are adding a zH to the callbacks of this Trainer, but there is already one. The currentzlist of callbacks is
:)r"   r&   r'   r   r   r   r   r*   )r,   r-   r   Zcb_classr/   r/   r0   r     s    
zCallbackHandler.add_callbackc                 C   sb   t |tr6| jD ]"}t ||r| j| |  S qn(| jD ] }||kr<| j| |  S q<d S r}   r"   r&   r   remover,   r-   r   r/   r/   r0   pop_callback  s    



zCallbackHandler.pop_callbackc                 C   sD   t |tr4| jD ] }t ||r| j|  d S qn| j| d S r}   r   r   r/   r/   r0   remove_callback  s    



zCallbackHandler.remove_callbackc                 C   s   d dd | jD S )Nr6   c                 s   s   | ]}|j jV  qd S r}   )r'   r(   r   r/   r/   r0   r     s     z0CallbackHandler.callback_list.<locals>.<genexpr>)joinr   rT   r/   r/   r0   r     s    zCallbackHandler.callback_listrh   c                 C   s   |  d|||S )Nrl   
call_eventr,   rU   r+   ri   r/   r/   r0   rl     s    zCallbackHandler.on_init_endc                 C   s   d|_ | d|||S )NFrm   )r^   r   r   r/   r/   r0   rm     s    zCallbackHandler.on_train_beginc                 C   s   |  d|||S )Nrn   r   r   r/   r/   r0   rn     s    zCallbackHandler.on_train_endc                 C   s   d|_ | d|||S )NFro   )r_   r   r   r/   r/   r0   ro     s    zCallbackHandler.on_epoch_beginc                 C   s   |  d|||S )Nrp   r   r   r/   r/   r0   rp     s    zCallbackHandler.on_epoch_endc                 C   s"   d|_ d|_d|_| d|||S )NFrq   )rb   ra   r`   r   r   r/   r/   r0   rq     s    zCallbackHandler.on_step_beginc                 C   s   |  d|||S )Nrr   r   r   r/   r/   r0   rr     s    z%CallbackHandler.on_pre_optimizer_stepc                 C   s   |  d|||S )Nrs   r   r   r/   r/   r0   rs     s    z!CallbackHandler.on_optimizer_stepc                 C   s   |  d|||S )Nrt   r   r   r/   r/   r0   rt     s    zCallbackHandler.on_substep_endc                 C   s   |  d|||S )Nru   r   r   r/   r/   r0   ru     s    zCallbackHandler.on_step_endc                 C   s   d|_ | jd||||dS )NFrv   rw   )ra   r   r,   rU   r+   ri   rw   r/   r/   r0   rv     s    zCallbackHandler.on_evaluatec                 C   s   | j d||||dS )Nrx   r   r   r   r/   r/   r0   rx     s    zCallbackHandler.on_predictc                 C   s   d|_ | d|||S )NFry   )r`   r   r   r/   r/   r0   ry     s    zCallbackHandler.on_savec                 C   s   d|_ | jd||||dS )NFrz   )logs)rb   r   )r,   rU   r+   ri   r   r/   r/   r0   rz     s    zCallbackHandler.on_logc                 C   s   |  d|||S )Nr{   r   r   r/   r/   r0   r{     s    z"CallbackHandler.on_prediction_stepc              
   K   sP   | j D ]D}t|||||f| j| j| j| j| j| jd|}|d k	r|}q|S )N)r   r   r   r   r   r   )r   getattrr   r   r   r   r   r   )r,   eventrU   r+   ri   rk   r-   resultr/   r/   r0   r     s$    

zCallbackHandler.call_eventN)r(   rH   rI   rJ   r   r   r   r   propertyr   r   r   r]   rl   rm   rn   ro   rp   rq   rr   rs   rt   ru   rv   rx   ry   rz   r{   r   r/   r/   r/   r0   r|     s.   	
r|   c                   @   s4   e Zd ZdZeeedddZeeedddZdS )r~   zx
    A [`TrainerCallback`] that handles the default flow of the training loop for logs, evaluation and checkpoints.
    rh   c                 K   s   |j dkr|jrd|_|jtjkr8|j |j dkr8d|_|jtjkrf|j |j dkrf|j	|j krfd|_
|jtjkr|jdkr|j |j dkrd|_|j |jkrd|_|jtjkrd|_|S )Nr   Tr   )r   Zlogging_first_steprb   logging_strategyr	   ZSTEPSr   eval_strategyr   
eval_delayra   save_strategyr   r`   r   r^   NOrj   r/   r/   r0   ru     s.    


zDefaultFlowCallback.on_step_endc                 K   sF   |j tjkrd|_|jtjkr0|j|jkr0d|_|jtjkrBd|_	|S )NT)
r   r	   EPOCHrb   r   r   r   ra   r   r`   rj   r/   r/   r0   rp   =  s    z DefaultFlowCallback.on_epoch_endN)	r(   rH   rI   rJ   r   r   r]   ru   rp   r/   r/   r/   r0   r~     s    r~   c                   @   sT   e Zd ZdZdd Zdd Zdd Zdd	d
Zdd Zdd Z	dddZ
dd ZdS )ProgressCallbackzU
    A [`TrainerCallback`] that displays the progress of training or evaluation.
    c                 C   s   d | _ d | _d S r}   )training_barprediction_barrT   r/   r/   r0   r   R  s    zProgressCallback.__init__c                 K   s    |j rt|jdd| _d| _d S )NT)totaldynamic_ncolsr   )r   r   r   r   current_steprj   r/   r/   r0   rm   V  s    zProgressCallback.on_train_beginc                 K   s&   |j r"| j|j| j  |j| _d S r}   )r   r   updater   r   rj   r/   r/   r0   ru   [  s    zProgressCallback.on_step_endNc                 K   sB   |j r>t|r>| jd kr2tt|| jd kdd| _| jd d S )NT)r   Zleaver   r   )r   r
   r   r   lenr   r   )r,   rU   r+   ri   r   rk   r/   r/   r0   r{   `  s    
  z#ProgressCallback.on_prediction_stepc                 K   s$   |j r | jd k	r| j  d | _d S r}   r   r   closerj   r/   r/   r0   rv   h  s    

zProgressCallback.on_evaluatec                 K   s$   |j r | jd k	r| j  d | _d S r}   r   rj   r/   r/   r0   rx   n  s    

zProgressCallback.on_predictc           
      K   sh   |j rd| jd k	rdi }| D ]\}}|||< q|dd }	d|krTt|d d|d< | jt| d S )Nr   r   r3   )r   r   rW   poproundr?   rN   )
r,   rU   r+   ri   r   rk   Zshallow_logsrZ   r[   _r/   r/   r0   rz   t  s    
zProgressCallback.on_logc                 K   s   |j r| j  d | _d S r}   )r   r   r   rj   r/   r/   r0   rn     s    
zProgressCallback.on_train_end)N)N)r(   rH   rI   rJ   r   rm   ru   r{   rv   rx   rz   rn   r/   r/   r/   r0   r   M  s   

r   c                   @   s   e Zd ZdZdddZdS )PrinterCallbackz?
    A bare [`TrainerCallback`] that just prints the logs.
    Nc                 K   s   | dd }|jrt| d S )Nr   )r   r   print)r,   rU   r+   ri   r   rk   r   r/   r/   r0   rz     s    zPrinterCallback.on_log)N)r(   rH   rI   rJ   rz   r/   r/   r/   r0   r     s   r   c                   @   sL   e Zd ZdZdeee dddZdd Zd	d
 Z	dd Z
edddZdS )EarlyStoppingCallbacka1  
    A [`TrainerCallback`] that handles early stopping.

    Args:
        early_stopping_patience (`int`):
            Use with `metric_for_best_model` to stop training when the specified metric worsens for
            `early_stopping_patience` evaluation calls.
        early_stopping_threshold(`float`, *optional*):
            Use with TrainingArguments `metric_for_best_model` and `early_stopping_patience` to denote how much the
            specified metric must improve to satisfy early stopping conditions. `

    This callback depends on [`TrainingArguments`] argument *load_best_model_at_end* functionality to set best_metric
    in [`TrainerState`]. Note that if the [`TrainingArguments`] argument *save_steps* differs from *eval_steps*, the
    early stopping will not occur until the next save step.
    r           early_stopping_patienceearly_stopping_thresholdc                 C   s   || _ || _d| _d S )Nr   r   r   early_stopping_patience_counter)r,   r   r   r/   r/   r0   r     s    zEarlyStoppingCallback.__init__c                 C   sV   |j rtjntj}|jd ks<|||jrDt||j | jkrDd| _n|  jd7  _d S )Nr   r   )Zgreater_is_betternpZgreaterZlessr   absr   r   )r,   rU   r+   ri   metric_valueoperatorr/   r/   r0   check_metric_value  s    

z(EarlyStoppingCallback.check_metric_valuec                 K   s8   |j std|jd k	s td|jtjks4tdd S )Nz<EarlyStoppingCallback requires load_best_model_at_end = Truez?EarlyStoppingCallback requires metric_for_best_model is definedzAEarlyStoppingCallback requires IntervalStrategy of steps or epoch)Zload_best_model_at_endAssertionErrormetric_for_best_modelr   r	   r   rj   r/   r/   r0   rm     s    
z$EarlyStoppingCallback.on_train_beginc                 K   sh   |j }|dsd| }||}|d krBtd| d d S | |||| | j| jkrdd|_d S )NZeval_z@early stopping required metric_for_best_model, but did not find z so early stopping is disabledT)	r   
startswithgetr   r   r   r   r   r^   )r,   rU   r+   ri   rw   rk   Zmetric_to_checkr   r/   r/   r0   rv     s    



z!EarlyStoppingCallback.on_evaluaterQ   c                 C   s   | j | jdd| jidS )Nr   r   rg   r   rT   r/   r/   r0   r+     s     zEarlyStoppingCallback.stateN)r   r   )r(   rH   rI   rJ   rM   r   rK   r   r   rm   rv   r#   r+   r/   r/   r/   r0   r     s   	r   ) rJ   r=   r;   r   typingr   r   r   r   Znumpyr   Z	tqdm.autor   Ztrainer_utilsr	   r
   Ztraining_argsr   utilsr   Z
get_loggerr(   r   r   r$   r]   r    r|   r~   r   r   r   r/   r/   r/   r0   <module>   s.   
u,=  5: