Ë
    ÿÍ:jq/  ã                   óÈ   — 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
 d dlmZmZ erd dlmZ  G d„ d«      Z	 	 	 	 dd	eeed
f      dee   dee   ded   ddf
d„Zdd„Zdd„Zy)é    )ÚTYPE_CHECKINGÚLiteralÚOptionalÚUnionN)ÚCallback)ÚMisconfigurationException)ÚEVAL_DATALOADERSÚTRAIN_DATALOADERS)Ú	_LRFinderc                   ó.  — e Zd ZdZdd„Z	 	 	 	 	 	 	 	 	 	 	 	 ddddeeedf      d	ee   d
ee   ded   de	d   de
dededede
dededee   fd„Z	 	 	 	 	 	 	 	 	 	 	 	 d dddeeedf      d	ee   d
ee   ded   de	d   dededede
dee   dede
ded   fd„Zy)!ÚTunerzTuner class to tune your model.ÚreturnNc                 ó   — || _         y )N)Ú_trainer)ÚselfÚtrainers     ús/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/pytorch_lightning/tuner/tuning.pyÚ__init__zTuner.__init__   s	   € Øˆ�ó    Úmodelzpl.LightningModuleÚtrain_dataloadersúpl.LightningDataModuleÚval_dataloadersÚdataloadersÚ
datamoduleÚmethod©ÚfitÚvalidateÚtestÚpredictÚmodeÚsteps_per_trialÚinit_valÚ
max_trialsÚbatch_arg_nameÚmarginÚmax_valc           	      ó°  — t        ||||«       t        | j                  «       d|cxk  rdk  sn J d|›�«       ‚ddlm}  ||||	|
|||¬«      }d|_        |g| j                  j                  z   | j                  _        |dk(  r| j                  j                  ||||«       nk|d	k(  r| j                  j                  |||¬
«       nG|dk(  r| j                  j                  |||¬
«       n#|dk(  r| j                  j                  |||¬
«       | j                  j                  D �cg c]	  }||usŒ|‘Œ c}| j                  _        |j                  S c c}w )ad
  Iteratively try to find the largest batch size for a given model that does not give an out of memory (OOM)
        error.

        Args:
            model: Model to tune.
            train_dataloaders: A collection of :class:`torch.utils.data.DataLoader` or a
                :class:`~pytorch_lightning.core.datamodule.LightningDataModule` specifying training samples.
                In the case of multiple dataloaders, please see this :ref:`section <multiple-dataloaders>`.
            val_dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them specifying validation samples.
            dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them specifying val/test/predict
                samples used for running tuner on validation/testing/prediction.
            datamodule: An instance of :class:`~pytorch_lightning.core.datamodule.LightningDataModule`.
            method: Method to run tuner on. It can be any of ``("fit", "validate", "test", "predict")``.
            mode: Search strategy to update the batch size:

                - ``'power'``: Keep multiplying the batch size by 2, until we get an OOM error.
                - ``'binsearch'``: Initially keep multiplying by 2 and after encountering an OOM error
                    do a binary search between the last successful batch size and the batch size that failed.

            steps_per_trial: number of steps to run with a given batch size.
                Ideally 1 should be enough to test if an OOM error occurs,
                however in practise a few are needed
            init_val: initial batch size to start the search with
            max_trials: max number of increases in batch size done before
               algorithm is terminated
            batch_arg_name: name of the attribute that stores the batch size.
                It is expected that the user has provided a model or datamodule that has a hyperparameter
                with that name. We will look for this attribute name in the following places

                - ``model``
                - ``model.hparams``
                - ``trainer.datamodule`` (the datamodule passed to the tune method)

            margin: Margin to reduce the found batch size by to provide a safety buffer. Only applied when using
                'binsearch' mode. Should be a float between 0 and 1. Defaults to 0.05 (5% reduction).
            max_val: Maximum batch size limit, defaults to 8192.
                Helps prevent testing unrealistically large or inefficient batch sizes (e.g., 2**25)
                when running on CPU or when automatic OOM detection is not available.

        g        g      ð?z1`margin` should be between 0 and 1. Found margin=r   ©ÚBatchSizeFinder)r"   r#   r$   r%   r&   r'   r(   Tr   r   )r   r    r!   )Ú_check_tuner_configurationÚ%_check_scale_batch_size_configurationr   Ú-pytorch_lightning.callbacks.batch_size_finderr+   Ú_early_exitÚ	callbacksr   r   r    r!   Úoptimal_batch_size)r   r   r   r   r   r   r   r"   r#   r$   r%   r&   r'   r(   r+   Úbatch_size_finderÚcbs                    r   Úscale_batch_sizezTuner.scale_batch_size   sP  € ôp 	#Ð#4°oÀ{ÐTZÔ[Ü-¨d¯m©mÔ<Ø�fÔ"˜sÔ"ÐZÐ&XÐQWÐPYÐ$ZÓZÐ"õ 	Rá&5ØØ+ØØ!Ø)ØØô'
Ðð )-ÐÔ%Ø#4Ð"5¸¿¹×8OÑ8OÑ"Oˆ�‰Ôà�UŠ?Ø�M‰M×Ñ˜eÐ%6¸ÈÕTØ�zÒ!Ø�M‰M×"Ñ" 5¨+À*Ð"ÕMØ�vÒØ�M‰M×Ñ˜u k¸jÐÕIØ�yÒ Ø�M‰M×!Ñ! %¨ÀÐ!ÔLà04·±×0GÑ0GÖ"g¨"È2ÐUfÒKf¢2Ò"gˆ�‰ÔØ ×3Ñ3Ð3ùò #hs   Ä'	EÄ1EÚmin_lrÚmax_lrÚnum_trainingÚearly_stop_thresholdÚupdate_attrÚ	attr_namer   c           	      óÀ  — |dk7  rt        d«      ‚t        ||||«       t        | j                  «       ddlm}  ||||	|
|||¬«      }d|_        |g| j                  j                  z   | j                  _        | j                  j                  ||||«       | j                  j                  D �cg c]	  }||usŒ|‘Œ c}| j                  _        |j                  S c c}w )aX  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:
            model: Model to tune.
            train_dataloaders: A collection of :class:`torch.utils.data.DataLoader` or a
                :class:`~pytorch_lightning.core.datamodule.LightningDataModule` specifying training samples.
                In the case of multiple dataloaders, please see this :ref:`section <multiple-dataloaders>`.
            val_dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them specifying validation samples.
            dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them specifying val/test/predict
                samples used for running tuner on validation/testing/prediction.
            datamodule: An instance of :class:`~pytorch_lightning.core.datamodule.LightningDataModule`.
            method: Method to run tuner on. It can be any of ``("fit", "validate", "test", "predict")``.
            min_lr: minimum learning rate to investigate
            max_lr: maximum learning rate to investigate
            num_training: number of learning rates to test
            mode: Search strategy to update learning rate after each batch:

                - ``'exponential'``: 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.
            attr_name: Name of the attribute which stores the learning rate. The names 'learning_rate' or 'lr' get
                automatically detected. Otherwise, set the name here.

        Raises:
            MisconfigurationException:
                If learning rate/lr in ``model`` or ``model.hparams`` isn't overridden,
                or if you are using more than one optimizer.

        r   z>method='fit' is the only valid configuration to run lr finder.r   ©ÚLearningRateFinder)r5   r6   Únum_training_stepsr"   r8   r9   r:   T)
r   r,   Ú_check_lr_find_configurationr   Ú%pytorch_lightning.callbacks.lr_finderr=   r/   r0   r   Ú
optimal_lr)r   r   r   r   r   r   r   r5   r6   r7   r"   r8   r9   r:   r=   Úlr_finder_callbackr3   s                    r   Úlr_findzTuner.lr_findw   s×   € ðd �UŠ?Ü+Ð,lÓmÐmä"Ð#4°oÀ{ÐTZÔ[Ü$ T§]¡]Ô3õ 	Má'9ØØØ+ØØ!5Ø#Øô(
Ðð *.ÐÔ&Ø#5Ð"6¸¿¹×9PÑ9PÑ"Pˆ�‰Ôà�‰×Ñ˜%Ð!2°OÀZÔPà04·±×0GÑ0GÖ"h¨"È2ÐUgÒKg¢2Ò"hˆ�‰Ôà!×,Ñ,Ð,ùò #is   Â/	CÂ9C©r   z
pl.Trainerr   N)NNNNr   Úpoweré   é   é   Ú
batch_sizegš™™™™™©?i    )NNNNr   g:Œ0âŽyE>é   éd   Úexponentialg      @TÚ )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r   r   r
   r	   r   ÚstrÚintÚfloatr4   ÚboolrC   © r   r   r   r      sö  „ Ù)ó ð [_Ø6:Ø26Ø9=Ø@EØØ ØØØ*ØØñV4à#ðV4ð $ EÐ*;Ð=UÐ*UÑ$VÑWðV4ð "Ð"2Ñ3ð	V4ð
 Ð.Ñ/ðV4ð Ð5Ñ6ðV4ð Ð<Ñ=ðV4ð ðV4ð ðV4ð ðV4ð ðV4ð ðV4ð ðV4ð ðV4ð 
�#‰óV4ðv [_Ø6:Ø26Ø9=Ø@EØØØØ!Ø03Ø ØñL-à#ðL-ð $ EÐ*;Ð=UÐ*UÑ$VÑWðL-ð "Ð"2Ñ3ð	L-ð
 Ð.Ñ/ðL-ð Ð5Ñ6ðL-ð Ð<Ñ=ðL-ð ðL-ð ðL-ð ðL-ð ðL-ð ' u™oðL-ð ðL-ð ðL-ð 
�+Ñ	ôL-r   r   r   r   r   r   r   r   r   c                 óˆ   — d}||vrt        d|›d|› d�«      ‚|dk(  r|�t        d|›d�«      ‚y | €|�t        d|›d	�«      ‚y )
Nr   zmethod z is invalid. Should be one of ú.r   zIn tuner with method=zs, `dataloaders` argument should be None, please consider setting `train_dataloaders` and `val_dataloaders` instead.zIn tuner with `method`=zt, `train_dataloaders` and `val_dataloaders` arguments should be None, please consider setting `dataloaders` instead.)Ú
ValueErrorr   )r   r   r   r   Úsupported_methodss        r   r,   r,   Æ   s–   € ð ?ÐØÐ&Ñ&Ü˜7 6 *Ð,JÐK\ÐJ]Ð]^Ð_Ó`Ð`à�‚ØÐ"Ü+Ø'¨ zð 2^ð ^óð ð #ð Ð(¨OÐ,GÜ+Ø)¨&¨ð 4\ð \óð ð -Hr   c                 ó€   — ddl m} | j                  D �cg c]  }t        ||«      sŒ|‘Œ }}|rt	        d«      ‚y c c}w )Nr   r<   zqTrainer is already configured with a `LearningRateFinder` callback.Please remove it if you want to use the Tuner.)r@   r=   r0   Ú
isinstancerY   )r   r=   r3   Úconfigured_callbackss       r   r?   r?   Þ   sI   € åHà)0×):Ñ):Öa 2¼jÈÐM_Õ>`šBÐaÐÐaÙÜð=ó
ð 	
ð ùò bs   •;§;c                 óÂ   — | j                   j                  rt        d«      ‚ddlm} | j
                  D �cg c]  }t        ||«      sŒ|‘Œ }}|rt        d«      ‚y c c}w )NzMTuning the batch size is currently not supported with distributed strategies.r   r*   znTrainer is already configured with a `BatchSizeFinder` callback.Please remove it if you want to use the Tuner.)Ú_accelerator_connectorÚis_distributedrY   r.   r+   r0   r\   )r   r+   r3   r]   s       r   r-   r-   ê   sf   € Ø×%Ñ%×4Ò4ÜÐhÓiÐiõ Nà)0×):Ñ):Ö^ 2¼jÈÈ_Õ>]šBÐ^ÐÐ^ÙÜð=ó
ð 	
ð ùò _s   ¶AÁA)NNNr   rD   )Útypingr   r   r   r   Úpytorch_lightningÚplÚ$pytorch_lightning.callbacks.callbackr   Ú&pytorch_lightning.utilities.exceptionsr   Ú!pytorch_lightning.utilities.typesr	   r
   Ú!pytorch_lightning.tuner.lr_finderr   r   r,   r?   r-   rV   r   r   ú<module>rh      sœ   ð÷ ;Ó :ã Ý 9Ý Lß QáÝ;÷j-ñ j-ð\ W[Ø26Ø.2Ø<Añ	Ø Ð&7Ð9QÐ&QÑ RÑSðàÐ.Ñ/ðð Ð*Ñ+ðð Ð8Ñ9ð	ð
 
óó0	
ô
r   