a
    da                     @   s`   d Z ddl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 G dd	 d	eZdS )
zD
LearningRateFinder
==================

Finds optimal learning rate
    )OptionalN)Callback)	_LRFinderlr_find)_TunerExitException)isolate_rngc                	   @   sX   e Zd ZdZdZdeeeeee e	d	d
ddZ
ddd	dddZddd	dddZd	S )LearningRateFindera  The ``LearningRateFinder`` callback enables the user to do a range test of good initial learning rates, to
    reduce the amount of guesswork in picking a good starting learning rate.

    Args:
        min_lr: Minimum learning rate to investigate

        max_lr: Maximum learning rate to investigate

        num_training_steps: Number of learning rates to test

        mode: Search strategy to update learning rate after each batch:

            - ``'exponential'`` (default): Increases the learning rate exponentially.
            - ``'linear'``: Increases the learning rate linearly.

        early_stop_threshold: Threshold for stopping the search. If the
            loss at any point is larger than early_stop_threshold*best_loss
            then the search is stopped. To disable, set to None.

        update_attr: Whether to update the learning rate attribute or not.

    Example::

        # Customize LearningRateFinder callback to run at different epochs.
        # This feature is useful while fine-tuning models.
        from pytorch_lightning.callbacks import LearningRateFinder


        class FineTuneLearningRateFinder(LearningRateFinder):
            def __init__(self, milestones, *args, **kwargs):
                super().__init__(*args, **kwargs)
                self.milestones = milestones

            def on_fit_start(self, *args, **kwargs):
                return

            def on_train_epoch_start(self, trainer, pl_module):
                if trainer.current_epoch in self.milestones or trainer.current_epoch == 0:
                    self.lr_find(trainer, pl_module)


        trainer = Trainer(callbacks=[FineTuneLearningRateFinder(milestones=(5, 10))])
        trainer.fit(...)

    Raises:
        MisconfigurationException:
            If learning rate/lr in ``model`` or ``model.hparams`` isn't overridden when ``auto_lr_find=True``,
            or if you are using more than one optimizer.
    )Zlinearexponential:0yE>   d   r	         @FN)min_lrmax_lrnum_training_stepsmodeearly_stop_thresholdupdate_attrreturnc                 C   sV   |  }|| jvr"td| j || _|| _|| _|| _|| _|| _d| _	d | _
d S )Nz`mode` should be either of F)lowerSUPPORTED_MODES
ValueError_min_lr_max_lr_num_training_steps_mode_early_stop_threshold_update_attr_early_exitZ	lr_finder)selfr   r   r   r   r   r    r    n/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/callbacks/lr_finder.py__init__R   s    	
zLearningRateFinder.__init__z
pl.Trainerzpl.LightningModule)trainer	pl_moduler   c                 C   s\   t  6 t||| j| j| j| j| j| jd| _W d    n1 sB0    Y  | j	rXt
 d S )N)r   r   Znum_trainingr   r   r   )r   r   r   r   r   r   r   r   Z
optimal_lrr   r   r   r#   r$   r    r    r!   r   i   s    &zLearningRateFinder.lr_findc                 C   s   |  || d S )N)r   r%   r    r    r!   on_fit_starty   s    zLearningRateFinder.on_fit_start)r
   r   r   r	   r   F)__name__
__module____qualname____doc__r   floatintstrr   boolr"   r   r&   r    r    r    r!   r      s&   2      r   )r*   typingr   Zpytorch_lightningplZ$pytorch_lightning.callbacks.callbackr   Z!pytorch_lightning.tuner.lr_finderr   r   Z&pytorch_lightning.utilities.exceptionsr   Z pytorch_lightning.utilities.seedr   r   r    r    r    r!   <module>   s   