a
    dUx                    @   s  d Z ddlZddlZddlZddlZddlZddlmZmZ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mZmZmZmZmZm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) ddl*m+Z+ ddl,m-Z- ddl.Z/ddl0m1Z1 ddl2m3Z3 ddl4m5Z5 ddl6m7Z7 ddl8m9Z9 ddl:m;Z;m<Z< ddl=m>Z>m?Z?m@Z@mAZA ddlBmCZC ddlDmEZE ddlFmGZG ddlHmIZI ddlJmKZKmLZL ddlMmNZN ddlOmPZP ddlQmRZRmSZS dd lTmUZUmVZVmWZWmXZX dd!lYmZZZ dd"l[m\Z\m]Z]m^Z^m_Z_m`Z` dd#lambZbmcZc dd$ldmeZe dd%lfmgZgmhZhmiZimjZj dd&lkmlZl dd'lmmnZn dd(lompZp dd)lqmrZr dd*lsmtZtmuZumvZv dd+lwmxZx dd,lymzZzm{Z{m|Z|m}Z} dd-l~mZ dd.lmZmZ dd/lmZmZ dd0lmZmZmZmZmZ dd1lmZ dd2lmZ dd3lmZmZ dd4lmZ dd5lmZ dd6lmZmZmZ dd7lmZ dd8lmZmZmZmZmZ eeZejd9d:d; G d<d= d=ZedBe;eed?d@dAZdS )Cz!Trainer to automate the training.    N)_ArgumentGroupArgumentParser	Namespace)contextmanager)deepcopy)	timedelta)Path)AnyDict	GeneratorIterableListOptionalTypeUnion)proxy)apply_to_collection)module_available)Version)Tensor)	Optimizer)
DataLoader)Literal)get_filesystem)_auto_add_worker_init_fn)_TORCH_GREATER_EQUAL_2_0)_PATH)PossibleUserWarning)AcceleratorTPUAccelerator)Callback
CheckpointEarlyStoppingProgressBarBase)BasePredictionWriter)LightningDataModule)Logger)TensorBoardLogger)PredictionLoopTrainingEpochLoop)EvaluationLoop)FitLoop)_parse_loop_limits_reset_progress)ApexMixedPrecisionPluginMixedPrecisionPluginPLUGIN_INPUTPrecisionPlugin)Profiler)DDPFullyShardedNativeStrategyDDPStrategyParallelStrategySingleDeviceStrategyStrategy)callsetup)verify_loop_configurations)_LITERAL_WARN_PRECISION_INPUT_PRECISION_INPUT_STRAcceleratorConnector)CallbackConnector)CheckpointConnector)DataConnector)LoggerConnector)	_OUT_DICT
_PBAR_DICT_ResultCollection)SignalConnector)RunningStage	TrainerFnTrainerStateTrainerStatus)CombinedLoader)_TunerResultTuner)GradClipAlgorithmTypeparsing)_defaults_from_env_varsadd_argparse_argsfrom_argparse_argsparse_argparserparse_env_variables)_add_capture_metadata_collate)has_len_all_ranks)ExitGracefullyExceptionMisconfigurationException)_fault_tolerant_training)is_overridden)rank_zero_deprecationrank_zero_inforank_zero_warn)isolate_rng)_EVALUATE_OUTPUT_PREDICT_OUTPUTEVAL_DATALOADERSLRSchedulerConfigTRAIN_DATALOADERSignorezXtorch.distributed.reduce_op is deprecated, please use torch.distributed.ReduceOp instead)messagec                7       sJ  e Zd Zedeeee ef eeee	e
 e
f  ee eeeef  ee eee eee	e eef  eee	e eef  ee eee	e eef  ee eeeef eeeef ee eeef eeeeeef f  ee ee eee eeeeeeef f  eeeef  eeeef  eeeef  eeeef  eeeef  eeeeef  eeeef  eeeeeeeef  eeeef  ee eeeef  eeeef eeeeef eeee	e f  ee ee eeedd4 fddZddddZeddddZddeeeef  ee ee ee ddddZddeeeef  ee ee ee ddddZ ded eeeef  ee eee e!dddZ"ded eeeef  ee eee eee#e!f  dddZ$ded eeeef  ee eee e!dd d!Z%ded eeeef  ee eee eee#e!f  dd"d#Z&ded eeeef  ee ee ee ee# d$d%d&Z'ded eeeef  ee ee ee ee# d$d'd(Z(ddeeeef  ee ee ee eeee)f  eeee)f  e*d* e+d+	d,d-Z,dee dd.d/d0Z-ddee eee!e#f  d1d2d3Z.ddd4d5Z/ddd6d7Z0eee#e!f  dd8d9Z1ddd:d;Z2ddd<d=Z3e!dd>d?Z4ee# dd@dAZ5dddBdCZ6dddDdEZ7dddFdGZ8dddHdIZ9ddJee)ed e)e)dKdLdMZ:ee)e)e)dNdOdPZ;ee)e)ddNdQdRZ<eee=f ddSdTZ>eee)f ddUdVdWZ?eee)f ddUdXdYZ@eee)f ddUdZd[ZAee)e)e)dNd\d]ZBeCedd^d_d`ZDdddadbZEded dddcddZFded dddedfZGded dddgdhZHded dddidjZIeJeddkdlZKeJeddmdnZLeJeMddodpZNeJeddqdrZOeJeddsdtZPeJeddudvZQeJeddwdxZReJeddydzZSeJe	e dd{d|ZTeJedd}d~ZUeJddddZVeJe	eW dddZXeXjYe	eW ddddZXeJe	eZ dddZ[eJe	e dddZ\e\jYe	e ddddZ\eJee dddZ]eJe^dddZ_eJee) dddZ`eJeeajbjc dddZdedjYeajbjcddddZdeJee dddZeeJedddZfeJeeee)f  dddZgeJedddZheJedddZieJedddZjeJeek dddZleJe	ek dddZmeJe	en dddZoeJeep dddZqeJe	ep dddZreJees dddZteJeeeef  dddZueJee dddZvdeeee) ddddZwexe=dddZyexe)eeze{f e)e)dddZ|exee{ezf ezdddZ}exezdddZ~exe{e)eee{f dddZeJedddZeJedddÄZejYeddĜddÄZeJedddǄZejYeddĜddǄZeJedddʄZejYeddĜddʄZeJeddd̈́ZejYeddĜdd̈́ZeJedddЄZejYeddĜddЄZeJedddӄZeJedddՄZejYeddĜddՄZeJeddd؄ZeJedddڄZeJee ddd܄ZeJee dddބZeJedddZeJee dddZeJedddZeJedddZejYeddddZeJedddZejYeddddZeJedddZejYeddddZeJedddZejYeddddZeJedddZeJeeeeef  dddZeJee dddZejYee ddddZeJe	e dddZejYee	e  ddddZeJedddZeJedd dZeJedddZeJee dddZddddZeddd	ZeJeeef dd
dZ  ZS (  TrainerTN           F2          r   max_size_cycle)4loggerenable_checkpointing	callbacksdefault_root_dirgradient_clip_valgradient_clip_algorithm	num_nodesnum_processesdevicesgpusauto_select_gpus	tpu_coresipusenable_progress_baroverfit_batchestrack_grad_normcheck_val_every_n_epochfast_dev_runaccumulate_grad_batches
max_epochs
min_epochs	max_steps	min_stepsmax_timelimit_train_batcheslimit_val_batcheslimit_test_batcheslimit_predict_batchesval_check_intervallog_every_n_stepsacceleratorstrategysync_batchnorm	precisionenable_model_summarynum_sanity_val_stepsresume_from_checkpointprofiler	benchmarkdeterministic!reload_dataloaders_every_n_epochsauto_lr_findreplace_sampler_ddpdetect_anomalyauto_scale_batch_sizepluginsamp_backend	amp_levelmove_metrics_to_cpumultiple_trainloader_modeinference_modereturnc4           6         sp  t    td t| jj dt   t	 | _
|durHt|}t| |2| _t||	|||| |
||!|'|+|(||"|/|0|.d| _t| | _t| | _t| |%| _t| | _t| | _t||d}4t||d}5|4j|5d |4| _t | _ t | _!t" | _#d| _$| j%|||||#|| |  | j%||)| |durRt&|t't(fsRt)d| d	|durt*+|, st-d
| dt*.  d	|dkrt&|t't(fs|dkrt(|dkst-d| d	|| _/|durt*|, nd| _0t(|| _1|3| _2|,| _3| 4  | j%|*|- t56| |& |  | j%|||1 |  |  |  |  |  |  |  t57| ||||||||$	 dS )a3  
        Customize every aspect of training via flags.

        Args:

            accelerator: Supports passing different accelerator types ("cpu", "gpu", "tpu", "ipu", "hpu", "mps", "auto")
                as well as custom accelerator instances.

            accumulate_grad_batches: Accumulates grads every k batches or as set up in the dict.
                Default: ``None``.

            amp_backend: The mixed precision backend to use ("native" or "apex").
                Default: ``'native''``.

                .. deprecated:: v1.9
                    Setting ``amp_backend`` inside the ``Trainer`` is deprecated in v1.8.0 and will be removed
                    in v2.0.0. This argument was only relevant for apex which is being removed.

            amp_level: The optimization level to use (O1, O2, etc...). By default it will be set to "O2"
                if ``amp_backend`` is set to "apex".

                .. deprecated:: v1.8
                    Setting ``amp_level`` inside the ``Trainer`` is deprecated in v1.8.0 and will be removed
                    in v2.0.0.

            auto_lr_find: If set to True, will make trainer.tune() run a learning rate finder,
                trying to optimize initial learning for faster convergence. trainer.tune() method will
                set the suggested learning rate in self.lr or self.learning_rate in the LightningModule.
                To use a different key set a string instead of True with the key name.
                Default: ``False``.

            auto_scale_batch_size: If set to True, will `initially` run a batch size
                finder trying to find the largest batch size that fits into memory.
                The result will be stored in self.batch_size in the LightningModule
                or LightningDataModule depending on your setup.
                Additionally, can be set to either `power` that estimates the batch size through
                a power search or `binsearch` that estimates the batch size through a binary search.
                Default: ``False``.

            auto_select_gpus: If enabled and ``gpus`` or ``devices`` is an integer, pick available
                gpus automatically. This is especially useful when
                GPUs are configured to be in "exclusive mode", such
                that only one process at a time can access them.
                Default: ``False``.

                .. deprecated:: v1.9
                    ``auto_select_gpus`` has been deprecated in v1.9.0 and will be removed in v2.0.0.
                    Please use the function :func:`~lightning_fabric.accelerators.cuda.find_usable_cuda_devices`
                    instead.

            benchmark: The value (``True`` or ``False``) to set ``torch.backends.cudnn.benchmark`` to.
                The value for ``torch.backends.cudnn.benchmark`` set in the current session will be used
                (``False`` if not manually set). If :paramref:`~pytorch_lightning.trainer.Trainer.deterministic` is set
                to ``True``, this will default to ``False``. Override to manually set a different value.
                Default: ``None``.

            callbacks: Add a callback or list of callbacks.
                Default: ``None``.

            enable_checkpointing: If ``True``, enable checkpointing.
                It will configure a default ModelCheckpoint callback if there is no user-defined ModelCheckpoint in
                :paramref:`~pytorch_lightning.trainer.trainer.Trainer.callbacks`.
                Default: ``True``.

            check_val_every_n_epoch: Perform a validation loop every after every `N` training epochs. If ``None``,
                validation will be done solely based on the number of training batches, requiring ``val_check_interval``
                to be an integer value.
                Default: ``1``.

            default_root_dir: Default path for logs and weights when no logger/ckpt_callback passed.
                Default: ``os.getcwd()``.
                Can be remote file paths such as `s3://mybucket/path` or 'hdfs://path/'

            detect_anomaly: Enable anomaly detection for the autograd engine.
                Default: ``False``.

            deterministic: If ``True``, sets whether PyTorch operations must use deterministic algorithms.
                Set to ``"warn"`` to use deterministic algorithms whenever possible, throwing warnings on operations
                that don't support deterministic mode (requires PyTorch 1.11+). If not set, defaults to ``False``.
                Default: ``None``.

            devices: Will be mapped to either `gpus`, `tpu_cores`, `num_processes` or `ipus`,
                based on the accelerator type.

            fast_dev_run: Runs n if set to ``n`` (int) else 1 if set to ``True`` batch(es)
                of train, val and test to find any bugs (ie: a sort of unit test).
                Default: ``False``.

            gpus: Number of GPUs to train on (int) or which GPUs to train on (list or str) applied per node
                Default: ``None``.

                .. deprecated:: v1.7
                    ``gpus`` has been deprecated in v1.7 and will be removed in v2.0.
                    Please use ``accelerator='gpu'`` and ``devices=x`` instead.

            gradient_clip_val: The value at which to clip gradients. Passing ``gradient_clip_val=None`` disables
                gradient clipping. If using Automatic Mixed Precision (AMP), the gradients will be unscaled before.
                Default: ``None``.

            gradient_clip_algorithm: The gradient clipping algorithm to use. Pass ``gradient_clip_algorithm="value"``
                to clip by value, and ``gradient_clip_algorithm="norm"`` to clip by norm. By default it will
                be set to ``"norm"``.

            limit_train_batches: How much of training dataset to check (float = fraction, int = num_batches).
                Default: ``1.0``.

            limit_val_batches: How much of validation dataset to check (float = fraction, int = num_batches).
                Default: ``1.0``.

            limit_test_batches: How much of test dataset to check (float = fraction, int = num_batches).
                Default: ``1.0``.

            limit_predict_batches: How much of prediction dataset to check (float = fraction, int = num_batches).
                Default: ``1.0``.

            logger: Logger (or iterable collection of loggers) for experiment tracking. A ``True`` value uses
                the default ``TensorBoardLogger`` if it is installed, otherwise ``CSVLogger``.
                ``False`` will disable logging. If multiple loggers are provided, local files
                (checkpoints, profiler traces, etc.) are saved in the ``log_dir`` of he first logger.
                Default: ``True``.

            log_every_n_steps: How often to log within steps.
                Default: ``50``.

            enable_progress_bar: Whether to enable to progress bar by default.
                Default: ``True``.

            profiler: To profile individual steps during training and assist in identifying bottlenecks.
                Default: ``None``.

            overfit_batches: Overfit a fraction of training/validation data (float) or a set number of batches (int).
                Default: ``0.0``.

            plugins: Plugins allow modification of core behavior like ddp and amp, and enable custom lightning plugins.
                Default: ``None``.

            precision: Double precision (64), full precision (32), half precision (16) or bfloat16 precision (bf16).
                Can be used on CPU, GPU, TPUs, HPUs or IPUs.
                Default: ``32``.

            max_epochs: Stop training once this number of epochs is reached. Disabled by default (None).
                If both max_epochs and max_steps are not specified, defaults to ``max_epochs = 1000``.
                To enable infinite training, set ``max_epochs = -1``.

            min_epochs: Force training for at least these many epochs. Disabled by default (None).

            max_steps: Stop training after this number of steps. Disabled by default (-1). If ``max_steps = -1``
                and ``max_epochs = None``, will default to ``max_epochs = 1000``. To enable infinite training, set
                ``max_epochs`` to ``-1``.

            min_steps: Force training for at least these number of steps. Disabled by default (``None``).

            max_time: Stop training after this amount of time has passed. Disabled by default (``None``).
                The time duration can be specified in the format DD:HH:MM:SS (days, hours, minutes seconds), as a
                :class:`datetime.timedelta`, or a dictionary with keys that will be passed to
                :class:`datetime.timedelta`.

            num_nodes: Number of GPU nodes for distributed training.
                Default: ``1``.

            num_processes: Number of processes for distributed training with ``accelerator="cpu"``.
                Default: ``1``.

                .. deprecated:: v1.7
                    ``num_processes`` has been deprecated in v1.7 and will be removed in v2.0.
                    Please use ``accelerator='cpu'`` and ``devices=x`` instead.

            num_sanity_val_steps: Sanity check runs n validation batches before starting the training routine.
                Set it to `-1` to run all batches in all validation dataloaders.
                Default: ``2``.

            reload_dataloaders_every_n_epochs: Set to a non-negative integer to reload dataloaders every n epochs.
                Default: ``0``.

            replace_sampler_ddp: Explicitly enables or disables sampler replacement. If not specified this
                will toggled automatically when DDP is used. By default it will add ``shuffle=True`` for
                train sampler and ``shuffle=False`` for val/test sampler. If you want to customize it,
                you can set ``replace_sampler_ddp=False`` and add your own distributed sampler.

            resume_from_checkpoint: Path/URL of the checkpoint from which training is resumed. If there is
                no checkpoint file at the path, an exception is raised. If resuming from mid-epoch checkpoint,
                training will start from the beginning of the next epoch.

                .. deprecated:: v1.5
                    ``resume_from_checkpoint`` is deprecated in v1.5 and will be removed in v2.0.
                    Please pass the path to ``Trainer.fit(..., ckpt_path=...)`` instead.

            strategy: Supports different training strategies with aliases
                as well custom strategies.
                Default: ``None``.

            sync_batchnorm: Synchronize batch norm layers between process groups/whole world.
                Default: ``False``.

            tpu_cores: How many TPU cores to train on (1 or 8) / Single TPU to train on (1)
                Default: ``None``.

                .. deprecated:: v1.7
                    ``tpu_cores`` has been deprecated in v1.7 and will be removed in v2.0.
                    Please use ``accelerator='tpu'`` and ``devices=x`` instead.

            ipus: How many IPUs to train on.
                Default: ``None``.

                .. deprecated:: v1.7
                    ``ipus`` has been deprecated in v1.7 and will be removed in v2.0.
                    Please use ``accelerator='ipu'`` and ``devices=x`` instead.

            track_grad_norm: -1 no tracking. Otherwise tracks that p-norm. May be set to 'inf' infinity-norm. If using
                Automatic Mixed Precision (AMP), the gradients will be unscaled before logging them.
                Default: ``-1``.

            val_check_interval: How often to check the validation set. Pass a ``float`` in the range [0.0, 1.0] to check
                after a fraction of the training epoch. Pass an ``int`` to check after a fixed number of training
                batches. An ``int`` value can only be higher than the number of training batches when
                ``check_val_every_n_epoch=None``, which validates after every ``N`` training batches
                across epochs or during iteration-based training.
                Default: ``1.0``.

            enable_model_summary: Whether to enable model summarization by default.
                Default: ``True``.

            move_metrics_to_cpu: Whether to force internal logged metrics to be moved to cpu.
                This can save some gpu memory, but can make training slower. Use with attention.
                Default: ``False``.

            multiple_trainloader_mode: How to loop over the datasets when there are multiple train loaders.
                In 'max_size_cycle' mode, the trainer ends one epoch when the largest dataset is traversed,
                and smaller datasets reload when running out of their data. In 'min_size' mode, all the datasets
                reload when reaching the minimum length of datasets.
                Default: ``"max_size_cycle"``.

            inference_mode: Whether to use :func:`torch.inference_mode` or :func:`torch.no_grad` during
                evaluation (``validate``/``test``/``predict``).
        initz(: Initializing trainer with parameters: N)ru   rv   ry   rz   r   r   rw   rt   r   r   r   r   rx   r   Zamp_typer   r   )r   r   )r   r   )
epoch_loopz5`gradient_clip_val` should be an int or a float. Got .z`gradient_clip_algorithm` z! is invalid. Allowed algorithms: ri   infr   zJ`track_grad_norm` must be a positive number or 'inf' (infinity norm). Got )8super__init__rf   _log_api_eventlogdetail	__class____name__localsrI   stateosfspathrA   _data_connectorr>   _accelerator_connectorrB   _logger_connectorr?   _callback_connectorr@   _checkpoint_connectorrF   _signal_connectorrM   tunerr+   r)   connectfit_loopr*   validate_loop	test_loopr(   predict_loop
_ckpt_pathZon_trainer_init
isinstanceintfloat	TypeErrorrN   Zsupported_typelowerrX   Zsupported_typesrr   rs   r}   _inference_mode_detect_anomaly_setup_on_initr9   Z_init_profilerZ_init_debugging_flags)6selfrn   ro   rp   rq   rr   rs   rt   ru   rv   rw   rx   ry   rz   r{   r|   r}   r~   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   Ztraining_epoch_loopr    j/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/pytorch_lightning/trainer/trainer.pyr   w   s      $










zTrainer.__init__)r   c                 C   sj   t |  d| _t | _td| _d | _g | _g | _	g | _
g | _d | _d | _d | _td| _td| _d S )NFr   z-inf)r9   Z_log_device_infoZshould_stoprI   r   r   num_training_batchestrain_dataloadernum_sanity_val_batchesnum_test_batchesnum_val_batchesnum_predict_batchestest_dataloadersval_dataloaderspredict_dataloaders_last_train_dl_reload_epoch_last_val_dl_reload_epochr   r   r   r   r   !  s    


zTrainer._setup_on_initzpl.LightningModule)modelr   c                 C   st   t s*t|tjs&tdt|j d|S ddlm} t||rJ|	|S t|tjrZ|S tdt|j dd S )Nz*`model` must be a `LightningModule`, got ``r   )OptimizedModulezM`model` must be a `LightningModule` or `torch._dynamo.OptimizedModule`, got `)
r   r   plZLightningModuler   type__qualname__Ztorch._dynamor   Zfrom_compiled)r   r   r   r   r   r   _maybe_unwrap_optimized5  s    

zTrainer._maybe_unwrap_optimized)r   train_dataloadersr   
datamodule	ckpt_pathr   c              	   C   s.   |  |}|| j_t| | j||||| dS )a  
        Runs the full optimization routine.

        Args:
            model: Model to fit.

            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.

            ckpt_path: Path/URL of the checkpoint from which training is resumed. Could also be one of two special
                keywords ``"last"`` and ``"hpc"``. If there is no checkpoint file at the path, an exception is raised.
                If resuming from mid-epoch checkpoint, training will start from the beginning of the next epoch.

            datamodule: An instance of :class:`~pytorch_lightning.core.datamodule.LightningDataModule`.
        N)r   r   _lightning_moduler8   _call_and_handle_interrupt	_fit_implr   r   r   r   r   r   r   r   r   fitD  s
    
zTrainer.fitc                 C   s   t d t| jj d tj| j_	t
j| j_d| _t|trJ|}d }|d usZ|d urj|d urjtd| jj||||d |p| j}| jj| jj	|d| jd ud| _| j|| jd | jjsJ d| _d S )	Nr   z: trainer fit stageTzXYou cannot pass `train_dataloader` or `val_dataloaders` to `trainer.fit(datamodule=...)`)r   r   r   model_providedZmodel_connectedr   F)rf   r   r   r   r   r   rH   FITTINGr   fnrJ   RUNNINGstatustrainingr   r%   rX   r   attach_datar   r   _set_ckpt_pathlightning_moduler   _runr   stoppedr   r   r   r   r   d  s4    




zTrainer._fit_impl)r   dataloadersr   verboser   r   c              	   C   sF   |du r| j du r.tdn| |}|| j_t| | j|||||S )a  
        Perform one evaluation epoch over the validation set.

        Args:
            model: The model to validate.

            dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them,
                or a :class:`~pytorch_lightning.core.datamodule.LightningDataModule` specifying validation samples.

            ckpt_path: Either ``"best"``, ``"last"``, ``"hpc"`` or path to the checkpoint you wish to validate.
                If ``None`` and the model instance was passed, use the current weights.
                Otherwise, the best model checkpoint from the previous ``trainer.fit`` call will be loaded
                if a checkpoint callback is configured.

            verbose: If True, prints the validation results.

            datamodule: An instance of :class:`~pytorch_lightning.core.datamodule.LightningDataModule`.

        Returns:
            List of dictionaries with metrics logged during the validation phase, e.g., in model- or callback hooks
            like :meth:`~pytorch_lightning.core.module.LightningModule.validation_step`,
            :meth:`~pytorch_lightning.core.module.LightningModule.validation_epoch_end`, etc.
            The length of the list corresponds to the number of validation dataloaders used.
        Nz^`Trainer.validate()` requires a `LightningModule` when it hasn't been passed in a previous run)r   r   r   r   r   r8   r   _validate_implr   r   r   r   r   r   r   r   r   validate  s     

zTrainer.validatec                 C   s   t d t| jj d tj| j_	t
j| j_d| _t|trJ|}d }|d ur^|r^td|d u rr| j}d}nd}|| j_| jj|||d | jj| jj	||| jd ud| _| j| _| j|| jd}| jjsJ d| _|S )	Nr   z: trainer validate stageTzHYou cannot pass both `trainer.validate(dataloaders=..., datamodule=...)`F)r   r   r   r   )rf   r   r   r   r   r   rH   
VALIDATINGr   r   rJ   r   r   
validatingr   r%   rX   r   r   r   r   r   r   r   r   r   Z_validated_ckpt_pathr   r   r   r   r   r   r   r   r   resultsr   r   r   r     s0    



zTrainer._validate_implc              	   C   sF   |du r| j du r.tdn| |}|| j_t| | j|||||S )a&  
        Perform one evaluation epoch over the test set.
        It's separated from fit to make sure you never run on your test set until you want to.

        Args:
            model: The model to test.

            dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them,
                or a :class:`~pytorch_lightning.core.datamodule.LightningDataModule` specifying test samples.

            ckpt_path: Either ``"best"``, ``"last"``, ``"hpc"`` or path to the checkpoint you wish to test.
                If ``None`` and the model instance was passed, use the current weights.
                Otherwise, the best model checkpoint from the previous ``trainer.fit`` call will be loaded
                if a checkpoint callback is configured.

            verbose: If True, prints the test results.

            datamodule: An instance of :class:`~pytorch_lightning.core.datamodule.LightningDataModule`.

        Returns:
            List of dictionaries with metrics logged during the test phase, e.g., in model- or callback hooks
            like :meth:`~pytorch_lightning.core.module.LightningModule.test_step`,
            :meth:`~pytorch_lightning.core.module.LightningModule.test_epoch_end`, etc.
            The length of the list corresponds to the number of test dataloaders used.
        NzZ`Trainer.test()` requires a `LightningModule` when it hasn't been passed in a previous run)r   r   r   r   r   r8   r   
_test_implr   r   r   r   test  s    !

zTrainer.testc                 C   s   t d t| jj d tj| j_	t
j| j_d| _t|trJ|}d }|d ur^|r^td|d u rr| j}d}nd}|| j_| jj|||d | jj| jj	||| jd ud| _| j| _| j|| jd}| jjsJ d| _|S )	Nr   z: trainer test stageTzDYou cannot pass both `trainer.test(dataloaders=..., datamodule=...)`F)r   r   r   r   )rf   r   r   r   r   r   rH   TESTINGr   r   rJ   r   r   testingr   r%   rX   r   r   r   r   r   r   r   r   r   Z_tested_ckpt_pathr   r   r   r   r   r   r     s0    



zTrainer._test_impl)r   r   r   return_predictionsr   r   c              	   C   sF   |du r| j du r.tdn| |}|| j_t| | j|||||S )as  
        Run inference on your data.
        This will call the model forward function to compute predictions. Useful to perform distributed
        and batched predictions. Logging is disabled in the predict hooks.

        Args:
            model: The model to predict with.

            dataloaders: A :class:`torch.utils.data.DataLoader` or a sequence of them,
                or a :class:`~pytorch_lightning.core.datamodule.LightningDataModule` specifying prediction samples.

            datamodule: The datamodule with a predict_dataloader method that returns one or more dataloaders.

            return_predictions: Whether to return predictions.
                ``True`` by default except when an accelerator that spawns processes is used (not supported).

            ckpt_path: Either ``"best"``, ``"last"``, ``"hpc"`` or path to the checkpoint you wish to predict.
                If ``None`` and the model instance was passed, use the current weights.
                Otherwise, the best model checkpoint from the previous ``trainer.fit`` call will be loaded
                if a checkpoint callback is configured.

        Returns:
            Returns a list of dictionaries, one for each provided dataloader containing their respective predictions.

        See :ref:`Lightning inference section<deploy/production_basic:Predict step with your LightningModule>` for more.
        Nz]`Trainer.predict()` requires a `LightningModule` when it hasn't been passed in a previous run)r   r   r   r   r   r8   r   _predict_impl)r   r   r   r   r   r   r   r   r   predictQ  s    "

zTrainer.predictc                 C   s   t d t| jj d tj| j_	t
j| j_d| _|| j_t|trR|}d }|d urf|rftd|d u rz| j}d}nd}| jj|||d | jj| jj	||| jd ud| _| j| _| j|| jd}| jjsJ d| _|S )	Nr  z: trainer predict stageTzGYou cannot pass both `trainer.predict(dataloaders=..., datamodule=...)`F)r   r   r   r   )rf   r   r   r   r   r   rH   
PREDICTINGr   r   rJ   r   r   
predictingr   r   r   r%   rX   r   r   r   r   r   r   r   Z_predicted_ckpt_pathr   r   )r   r   r   r   r   r   r   r   r   r   r   r     s0    



zTrainer._predict_implr   )r   r   r   r  )	r   r   r   r   r   scale_batch_size_kwargslr_find_kwargsmethodr   c	           
      C   sZ   |  |}td t , | jj||||||||d}	W d   n1 sL0    Y  |	S )a  
        Runs routines to tune hyperparameters before training.

        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`.

            scale_batch_size_kwargs: Arguments for :func:`~pytorch_lightning.tuner.batch_size_scaling.scale_batch_size`

            lr_find_kwargs: Arguments for :func:`~pytorch_lightning.tuner.lr_finder.lr_find`

            method: Method to run tuner on. It can be any of ``("fit", "validate", "test", "predict")``.
        tune)r  r  r  N)r   rf   r   r^   r   Z_tune)
r   r   r   r   r   r   r  r  r  resultr   r   r   r    s    "

$zTrainer.tune)checkpoint_pathr   c                 C   sF   | j | | j   | j   | j   | jjtjkrB| j 	  d S N)
r   Zresume_startZ_restore_quantization_callbacksZrestore_modelZrestore_datamoduler   r   rH   r   Zrestore_callbacks)r   r	  r   r   r   _restore_modules_and_callbacks  s    


z&Trainer._restore_modules_and_callbacks)r   r   r   c                    s  |j d urdtttg} jd urdt fdd|D sdddd |D }td jjj	 d| d j
jtjkrt j j j j \}}| j_| j_t|drt|j  j|  j   j  t  t jj	 d	  j   t jj	 d
  j!   "   #   jj$sTt jj	 d|   %| t jj	 d  &  dt'j( dt'j) dt'j* d jj! d j+ d j, d j- d j. d j. d  j/0   j/1   j2   j
jtjkr 3d  4d  5   jj$r8t jj	 d|   %| t jj	 d  j67   j68   + }t jj	 d  9   j
jtjkr 3d  4d t jj	 d  :  t;j< j
_=d  j
_>|S )Nc                 3   s   | ]}t  j|V  qd S r
  )r   r   .0sr   r   r   	<genexpr>      zTrainer._run.<locals>.<genexpr>z, c                 s   s   | ]}|j V  qd S r
  )r   r  r   r   r   r    r  zBUsing a compiled model is incompatible with the current strategy: z. Only zl support compilation. Either switch to one of the supported strategies or avoid passing in a compiled model.hparamsz: preparing dataz!: setting up strategy environmentz7: restoring module and callbacks from checkpoint path: z: configuring sharded modelz?
             Lightning internal flow looks like this:
        z or z  ||
                                |                             ||
                         spawn processes                      ||
                 a6              ||
                                |                             ||
                        setup accelerator                     ||
                           and strategy                       ||  LIGHTNING
                                |                             ||
                        zw                     ||  FLOW
                                |                             ||
                        z;                     ||  DIRECTION
                     or z-                  ||
                     or z                   ||
                                |                             ||
                             results                          \/
        This is used to guide readers to the core loops: train, test, predict.
        zM is the simplest to understand, use `Go to Definition` to read it :)
        Zon_fit_startz: restoring training statez: trainer tearing downZ
on_fit_endz: calling teardown hooks)?Z_compiler_ctxr6   r4   r3   r   anyjoinRuntimeErrorr   r   r   r   rH   r   r,   r   r   r   r   r   hasattrrO   Zclean_namespacer  r   r   Z_attach_model_callbacksZ_attach_model_logging_functionsr:   r   r   r   Zprepare_dataZsetup_environment_Trainer__setup_profiler_call_setup_hookZrestore_checkpoint_after_setupr  _call_configure_sharded_modelrf   r   r   r  
_run_stage
_run_train_run_evaluate_run_predictr   reset_resultsreset_metricsr9   _call_callback_hooks_call_lightning_module_hook_log_hyperparamsr   Zrestore_training_stateZ
resume_end	_teardown_call_teardown_hookrJ   FINISHEDr   stage)r   r   r   Zsupported_strategiesZsupported_strategy_namesr   r   r   r   r   r   r     s    

 


















zTrainer._runc           
      C   s<  | j s
d S d }| jd ur | jjnd}| jjr|r| jj}| jj}g }| | @ D ]j}|| ||  }}t|t|kr|| qTt|t	rt
|t
|kr|| qT||krT|| qT|rtd| di ||}n | jjr| jj}n|r| jj}| j D ].}	|d ur |	| |	| j |	  qd S )NFz&Error while merging hparams: the keys zg are present in both the LightningModule's and LightningDataModule's hparams but have different values.)loggersr   r!  r   hparams_initialkeysr   appendr   r   idrX   Zlog_hyperparamsZ	log_graphsave)
r   r'  Zdatamodule_log_hyperparamsZdatamodule_hparamsZlightning_hparamsZinconsistent_keyskeyZlm_valZdm_valrn   r   r   r   r!  m  s<    




zTrainer._log_hyperparamsc                 C   s8   | j   | j}|dur |  | j  | j  dS )zThis is the Trainer's internal teardown, unrelated to the `teardown` hooks in LightningModule and
        Callback; those are handled by :meth:`_call_teardown_hook`.N)r   teardown_active_loopr   r   r   loopr   r   r   r"    s    

zTrainer._teardownc                 C   s@   | j d | j |  | jr&|  S | jr4|  S |   d S )Nz	run-stage)r   barrierdispatch
evaluatingr  r  r  r  r   r   r   r   r    s    zTrainer._run_stagec                 C   s   | j d | j  d S )NZsetup_training)r   r1  r   Zregister_signal_handlersr   r   r   r   _pre_training_routine  s    zTrainer._pre_training_routinec                 C   s   |    t  |   W d    n1 s,0    Y  | jd usDJ | j  td | | j_tj	
| j | j  W d    n1 s0    Y  d S )NT)r4  r^   _run_sanity_checkr   ZtraintorchZset_grad_enabledr   trainerZautogradZset_detect_anomalyr   runr   r   r   r   r    s    &

zTrainer._run_trainc              	   C   s   | j s
J | j  | | j_| jd| jj dH t| j	| j
 | j }W d    n1 sd0    Y  W d    n1 s0    Y  |D ]:}t|tr| D ]"\}}t|tr|  ||< qq|S )NZrun_Z_evaluation)r3  _evaluation_loop_reload_evaluation_dataloadersr7  r   profiler   r%  _evaluation_contextr   r   r8  r   dictitemsr   cpuitem)r   Zeval_loop_resultsr  kvr   r   r   r    s    

F

zTrainer._run_evaluatec                 C   sP   |  | j | | j_t| j| j | j W  d    S 1 sB0    Y  d S r
  )reset_predict_dataloaderr   r   r7  r<  r   r   r8  r   r   r   r   r    s    zTrainer._run_predictc                    s    j jj} jo  jdko |j }|rΈ jj}d _ j	
   j	   d |   fdd jD  _t  |  W d    n1 s0    Y   d  j	
   j	  t| | j_d S )Nr   TZon_sanity_check_startc                    s   g | ]}t  j|qS r   )minr   )r  Zval_batchesr   r   r   
<listcomp>  s   z-Trainer._run_sanity_check.<locals>.<listcomp>Zon_sanity_check_end)r   r   val_loopenable_validationr   Z
restartingr   r%  sanity_checkingr   r  r  r  r:  r   r   r6  no_gradr8  r-   )r   rF  Zshould_sanity_checkr%  r   r   r   r5    s.    
	




&


zTrainer._run_sanity_checkc                 C   sh   | j jd usJ | j j}| jd | jd ur<| jd|d | jd|d | jd|d | jd d S )NZ	pre_setupr9   r%  Z
post_setup)r   r   r   r1  r   _call_lightning_datamodule_hookr  r   )r   r   r   r   r   r    s    
zTrainer._call_setup_hookc                 C   sV   | j  8 tdr*ddlm} || j | d W d    n1 sH0    Y  d S )Nztorchdistx.deferred_initr   )materialize_moduleZconfigure_sharded_model)r   Zmodel_sharded_contextr   Ztorchdistx.deferred_initrL  r   r   )r   rL  r   r   r   r    s
    
z%Trainer._call_configure_sharded_modelc                 C   s   | j jd usJ | j j}| jd ur0| jd|d | jd|d | jd|d d | j_d | j_| j	D ]}|
d qb| j  d S )Nr-  rJ  success)r   r   r   rK  r  r   r   _current_fx_nameZ_metric_attributesr&  finalizer   Zdescribe)r   r   rn   r   r   r   r#  "  s    

zTrainer._call_teardown_hook)	pl_module)	hook_nameargsrP  kwargsr   c                O   s   |p| j }|d u rtdt||}t|s0d S |j}||_| jd|jj d|  ||i |}W d    n1 s|0    Y  ||_|S )Nz3No `LightningModule` is available to call hooks on.z[LightningModule]r   )	r   r   getattrcallablerN  r   r;  r   r   )r   rQ  rP  rR  rS  r   prev_fx_nameoutputr   r   r   r   8  s    

,z#Trainer._call_lightning_module_hook)rQ  rR  rS  r   c                 O   sv   | j d u rtdt| j |}t|rr| jd| j jj d|  ||i |W  d    S 1 sh0    Y  d S )Nz7No `LightningDataModule` is available to call hooks on.z[LightningDataModule]r   )r   r   rT  rU  r   r;  r   r   )r   rQ  rR  rS  r   r   r   r   rK  S  s    
 z'Trainer._call_lightning_datamodule_hookc              	   O   s   t | jj d|  | j}|r.|j}||_| jD ]j}t||}t|r4| j	
d|j d| * || | jg|R i | W d    q41 s0    Y  q4|r||_d S )Nz: calling callback hook: 
[Callback]r   )r   debugr   r   r   rN  rp   rT  rU  r   r;  	state_key)r   rQ  rR  rS  rP  rV  callbackr   r   r   r   r  a  s    

:zTrainer._call_callback_hooksc                 C   s*   i }| j D ]}| }|r
|||j< q
|S )z~Called when saving a model checkpoint, calls and returns every callback's `state_dict`, keyed by
        `Callback.state_key`.)rp   
state_dictrZ  )r   Zcallback_state_dictsr[  r\  r   r   r   _call_callbacks_state_dictx  s    
z"Trainer._call_callbacks_state_dict)
checkpointr   c              	   C   s   | j }|r|j}d|_| jD ]f}| jd|j d  || | j |}W d   n1 s\0    Y  |durtd|jj	 dq|r||_dS )zXCalled when saving a model checkpoint, calls every callback's `on_save_checkpoint` hook.on_save_checkpointrX  z.on_save_checkpointNzReturning a value from `z.on_save_checkpoint` was deprecated in v1.6 and is no longer supported as of v1.8. Please override `Callback.state_dict` to return state to be saved.)
r   rN  rp   r   r;  rZ  r_  
ValueErrorr   r   )r   r^  rP  rV  r[  r   r   r   r   "_call_callbacks_on_save_checkpoint  s    
.z*Trainer._call_callbacks_on_save_checkpointc              	      s   | j }|r|j}d|_|d}|du r,dS t|d tdk   fdd| jD }| | }|rxtdt| d	 | jD ]J}| j	d
|j
 d  || | j | W d   q~1 s0    Y  q~|r||_dS )zCalled when loading a model checkpoint.

        Calls every callback's `on_load_checkpoint` hook. We have a dedicated function for this rather than using
        `_call_callback_hooks` because we have special logic for getting callback_states.
        on_load_checkpointrp   Nzpytorch-lightning_versionz1.5.0devc                    s   h | ]} r|j n|jqS r   )_legacy_state_keyrZ  r  cbZis_legacy_ckptr   r   	<setcomp>  r  z=Trainer._call_callbacks_on_load_checkpoint.<locals>.<setcomp>zBe aware that when using `ckpt_path`, callbacks used to create the checkpoint need to be provided during `Trainer` instantiation. Please add the following callbacks: r   rX  z.on_load_checkpoint)r   rN  getr   rp   r(  r]   listr   r;  rZ  rb  )r   r^  rP  rV  callback_statesZcurrent_callbacks_keys
differencer[  r   rf  r   "_call_callbacks_on_load_checkpoint  s*    

0z*Trainer._call_callbacks_on_load_checkpointc                 C   sR   | d}|du rdS | jD ]0}| |j| |j}|rt|}|| qdS )zQCalled when loading a model checkpoint, calls every callback's `load_state_dict`.rp   N)rh  rp   rZ  rc  r   Zload_state_dict)r   r^  rj  r[  r   r   r   r   _call_callbacks_load_state_dict  s    

z'Trainer._call_callbacks_load_state_dictc                 O   s   | j }|j}||_t| j|}t|s*d S | jd| jjj d|  ||i |}W d    n1 sl0    Y  ||_|S )Nz
[Strategy]r   )	r   rN  rT  r   rU  r   r;  r   r   )r   rQ  rR  rS  rP  rV  r   rW  r   r   r   _call_strategy_hook  s     ,zTrainer._call_strategy_hook)eventr   c                 C   s   t jd|   d S )Nzlightning.trainer.)r6  Z_CZ_log_api_usage_once)ro  r   r   r   r     s    zTrainer._log_api_eventc                 C   sN   | j jd usJ | jdkr | jnd }t| j| j_| jj| j j|| j	d d S )Nrg   )r%  
local_ranklog_dir)
r   r   
world_sizerp  r   r   r   r   r9   rq  )r   rp  r   r   r   Z__setup_profiler  s    zTrainer.__setup_profilerc           
   
   C   s  | j j}|p| j}td|}| jdk}| r6|r6|s:dS | j tj| _	| j
dkrj| j j| j	tjd| _	t| j	ttf| j jtjd| _	t| j	tr| j	jn| j	}t|t| j jd t|tt| jd t rt|tt t| j	tst|| j j| _	|p| jp| j}t| j	| j|r$t| j	ntd }| _|dkr@dS | j| _t| jt rft!|| j| _n6| jtdkrt || j | _n| jdkrt"d	t| j#t r| j#| _$| j$| jkr<| j%dur<t&d
| j# d| j dnTt| j	| j|s| j#dkrtd| _$nt"dn t | j| j# | _$t'd| j$| _$| j(rp| j| j)k rpt*d| j d| j) dt+d | jdkr| jdkrt| jtr|tdkrd| }	t"d| j d| j d| d|	 d	dS )zResets the train dataloader and initialises required variables (number of batches, when to validate,
        etc.).

        Args:
            model: The ``LightningModule`` if calling this outside of the trainer scope.
        Ztraining_stepr   N)moder   )Zrankr   g      ?zWhen using an `IterableDataset`, `Trainer(limit_train_batches)` must be `1.0` or an int.An int specifies `num_training_batches` to use.z`val_check_interval` (zD) must be less than or equal to the number of the training batches (z). If you want to disable validation set `limit_val_batches` to 0.0 instead.If you want to validate based on the total training batches, set `check_val_every_n_epoch=None`.zWhen using an IterableDataset for `train_dataloader`, `Trainer(val_check_interval)` must be `1.0` or an int. An int k specifies checking validation every k training batches.rg   z The number of training batches (zA) is smaller than the logging interval Trainer(log_every_n_steps=zZ). Set a lower value for log_every_n_steps if you want to see logs for the training epoch.)categoryrh   zYou requested to check z of the `train_dataloader` but z * z\ < 1. Please increase the `limit_train_batches` argument. Try at least `limit_train_batches=r   ),r   Z_train_dataloader_sourcer   rZ   r   
is_definedZ_request_dataloaderrG   TRAININGr   r|   Z_resolve_overfit_batchesr   r   rK   Z_prepare_dataloaderr   loadersZ_worker_checkr   global_rankrY   rU   r   r   rV   r   lenr   r   current_epochr   r   rD  rX   r   Zval_check_batchr~   r`  maxr&  r   r]   r   )
r   r   sourcerP  has_stepZenable_trainingrw  moduleZorig_train_batchesZmin_percentager   r   r   reset_train_dataloader  s    











zTrainer.reset_train_dataloaderc                 C   s   | j j}| jp|}td|}| jdk}| r~|r~|r~| jjtj	krd| j
rV| jj s\| j
sd| j| _| j jtj|d\| _| _dS )zResets the validation dataloader and determines the number of batches.

        Args:
            model: The ``LightningModule`` if called outside of the trainer scope.
        validation_stepr   r   N)r   _val_dataloader_sourcer   rZ   r   ru  r   r   rH   r   rH  r   r   Z_should_check_val_epochrz  r   _reset_eval_dataloaderrG   r   r   r   )r   r   r|  rP  r}  rG  r   r   r   reset_val_dataloader^  s     



zTrainer.reset_val_dataloaderc                 C   sT   | j j}| jp|}td|}| jdk}| rP|rP|rP| j jtj|d\| _	| _
dS )zResets the test dataloader and determines the number of batches.

        Args:
            model: The ``LightningModule`` if called outside of the trainer scope.
        Z	test_stepr   r  N)r   Z_test_dataloader_sourcer   rZ   r   ru  r  rG   r   r   r   )r   r   r|  rP  r}  Zenable_testingr   r   r   reset_test_dataloaderu  s    


zTrainer.reset_test_dataloaderc                 C   sF   | j j}| jp|}| jdk}| rB|rB| j jtj|d\| _| _	dS )zResets the predict dataloader and determines the number of batches.

        Args:
            model: The ``LightningModule`` if called outside of the trainer scope.
        r   r  N)
r   Z_predict_dataloader_sourcer   r   ru  r  rG   r  r   r   )r   r   r|  rP  Zenable_predictionr   r   r   rC    s    

z Trainer.reset_predict_dataloaderc                 C   s   | j jsJ | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.acceleratorc                 C   s   | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.strategyc                 C   s   | j jS r
  )r   precision_pluginr   r   r   r   r    s    zTrainer.precision_pluginc                 C   s   | j jS r
  )r   rx  r   r   r   r   rx    s    zTrainer.global_rankc                 C   s   t | jddS )Nrp  r   rT  r   r   r   r   r   rp    s    zTrainer.local_rankc                 C   s   t | jddS )N	node_rankr   r  r   r   r   r   r    s    zTrainer.node_rankc                 C   s   t | jddS )Nrr  rg   r  r   r   r   r   rr    s    zTrainer.world_sizec                 C   s   t | jddS )Nrt   rg   r  r   r   r   r   rt     s    zTrainer.num_nodesc                 C   sv   t | jtr| jjn| jjg}|dus*J g }t|D ]:\}}t |tjr\||j	pV| q6t |t
r6|| q6|S )z List of device indexes per node.N)r   r   r5   Zparallel_devicesroot_device	enumerater6  devicer)  indexr   )r   rv   
device_idsidxr  r   r   r   r    s    


zTrainer.device_idsc                 C   s
   t | jS )z,Number of devices the trainer uses per node.)ry  r  r   r   r   r   num_devices  s    zTrainer.num_devicesc                 C   s   | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.lightning_modulec                 C   s   | j jS r
  r   
optimizersr   r   r   r   r    s    zTrainer.optimizers)
new_optimsr   c                 C   s   || j _d S r
  r  )r   r  r   r   r   r    s    c                 C   s   | j jS r
  )r   lr_scheduler_configsr   r   r   r   r    s    zTrainer.lr_scheduler_configsc                 C   s   | j jS r
  r   optimizer_frequenciesr   r   r   r   r    s    zTrainer.optimizer_frequencies)	new_freqsr   c                 C   s   || j _d S r
  r  )r   r  r   r   r   r    s    c                 C   s0   t ddd t| jtrdS t| jtr,dS d S )Na  The NVIDIA/apex AMP implementation has been deprecated upstream. Consequently, its integration inside PyTorch Lightning has been deprecated in v1.9.0 and will be removed in v2.0.0. Accessing `Trainer.amp_backend` will not be supported. You can assume it will be `'native'`   
stacklevelZapexZnative)r[   r   r  r.   r/   r   r   r   r   r     s    zTrainer.amp_backendc                 C   s
   | j jjS r
  )r   r  r   r   r   r   r   r     s    zTrainer.precisionc                 C   s   t | jdd S )Nscaler)rT  r  r   r   r   r   r    s    zTrainer.scalerc                 C   s   | j jS )zThe LightningModule, but possibly wrapped into DataParallel or DistributedDataParallel.

        To access the pure LightningModule, use
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.lightning_module` instead.
        r   r   r   r   r   r   r      s    zTrainer.modelc                 C   s   || j _dS )aS  Setter for the model, pass-through to accelerator and plugin where the model reference is stored. Used
        by the Tuner to reset the state of Trainer and Accelerator.

        Args:
            model: The LightningModule, possibly wrapped into DataParallel or DistributedDataParallel, depending
                on the backend.
        Nr  )r   r   r   r   r   r   	  s    	c                 C   sP   t | jdkr:t| jd ts,| jd j}q@| jd j}n| j}| j|}|S Nr   )	ry  r&  r   r'   save_dirrq  rq   r   	broadcast)r   dirpathr   r   r   rq    s    zTrainer.log_dirc                 C   s   | j jS r
  )r   is_global_zeror   r   r   r   r  %  s    zTrainer.is_global_zeroc                 C   s   t | jtr| jjS d S r
  )r   r   r5   distributed_sampler_kwargsr   r   r   r   r  )  s    z"Trainer.distributed_sampler_kwargsc                 C   s   t | jtS r
  )r   r   r5   r   r   r   r   data_parallel.  s    zTrainer.data_parallelc                 C   s"   | j j o td| jo | jdkS )z2Check if we should run validation during training.r  r   )r   r  ru  rZ   r   r   r   r   r   r   rG  2  s
    
zTrainer.enable_validationc                 C   s$   t | jjdkrtj| jS | jS )zThe default location to save artifacts of loggers, checkpoints etc.

        It is used as a fallback if logger or checkpoint callback do not define specific save paths.
        file)r   Z_default_root_dirprotocolr   pathnormpathr   r   r   r   rq   ;  s    zTrainer.default_root_dirc                 C   s   | j }t|dkr|d S dS )zThe first :class:`~pytorch_lightning.callbacks.early_stopping.EarlyStopping` callback in the
        Trainer.callbacks list, or ``None`` if it doesn't exist.r   N)early_stopping_callbacksry  r   rp   r   r   r   early_stopping_callbackE  s    zTrainer.early_stopping_callbackc                 C   s   dd | j D S )zA list of all instances of :class:`~pytorch_lightning.callbacks.early_stopping.EarlyStopping` found in
        the Trainer.callbacks list.c                 S   s   g | ]}t |tr|qS r   )r   r"   r  cr   r   r   rE  P  r  z4Trainer.early_stopping_callbacks.<locals>.<listcomp>rp   r   r   r   r   r  L  s    z Trainer.early_stopping_callbacksc                 C   s   dd | j D S )zA list of all instances of :class:`~pytorch_lightning.callbacks.prediction_writer.BasePredictionWriter`
        found in the Trainer.callbacks list.c                 S   s   g | ]}t |tr|qS r   )r   r$   rd  r   r   r   rE  V  r  z7Trainer.prediction_writer_callbacks.<locals>.<listcomp>r  r   r   r   r   prediction_writer_callbacksR  s    z#Trainer.prediction_writer_callbacksc                 C   s   | j }t|dkr|d S dS )zThe first :class:`~pytorch_lightning.callbacks.model_checkpoint.ModelCheckpoint` callback in the
        Trainer.callbacks list, or ``None`` if it doesn't exist.r   N)checkpoint_callbacksry  r  r   r   r   checkpoint_callbackX  s    zTrainer.checkpoint_callbackc                 C   s   dd | j D S )zA list of all instances of :class:`~pytorch_lightning.callbacks.model_checkpoint.ModelCheckpoint` found
        in the Trainer.callbacks list.c                 S   s   g | ]}t |tr|qS r   )r   r!   r  r   r   r   rE  c  r  z0Trainer.checkpoint_callbacks.<locals>.<listcomp>r  r   r   r   r   r  _  s    zTrainer.checkpoint_callbacksc                 C   s"   | j D ]}t|tr|  S qdS )zAn instance of :class:`~pytorch_lightning.callbacks.progress.base.ProgressBarBase` found in the
        Trainer.callbacks list, or ``None`` if one doesn't exist.N)rp   r   r#   )r   r  r   r   r   progress_bar_callbacke  s    


zTrainer.progress_bar_callbackc                 C   s    | j j}|d urtddd |S )Nz`trainer.resume_from_checkpoint` is deprecated in v1.5 and will be removed in v2.0. Specify the fit checkpoint path with `trainer.fit(ckpt_path=)` instead.   r  )r   Zresume_from_checkpoint_fit_pathr[   )r   r   r   r   r   r   n  s    zTrainer.resume_from_checkpointc                 C   s   | j S )aG  Set to the path/URL of a checkpoint loaded via :meth:`~pytorch_lightning.trainer.trainer.Trainer.fit`,
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.validate`,
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.test`, or
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.predict`. ``None`` otherwise.)r   r   r   r   r   r   z  s    zTrainer.ckpt_path)filepathweights_onlystorage_optionsr   c                 C   s(   | j du rtd| jj|||d dS )a*  
        Runs routine to create a checkpoint.

        Args:
            filepath: Path where checkpoint is saved.
            weights_only: If ``True``, will only save the model weights.
            storage_options: parameter for how to save to storage, passed to ``CheckpointIO`` plugin

        NzSaving a checkpoint is only possible if a model is attached to the Trainer. Did you call `Trainer.save_checkpoint()` before calling `Trainer.{fit,validate,test,predict}`?)r  r  )r   AttributeErrorr   save_checkpoint)r   r  r  r  r   r   r   r    s
    
zTrainer.save_checkpointc                 C   s   t | }dd |j D S )Nc                 S   s   i | ]\}}||j qS r   )default)r  rA  rB  r   r   r   
<dictcomp>  r  z.Trainer.default_attributes.<locals>.<dictcomp>)inspect	signature
parametersr>  )clsZinit_signaturer   r   r   default_attributes  s    
zTrainer.default_attributes)r  rR  rS  r   c                 K   s   t | |fi |S r
  )rR   )r  rR  rS  r   r   r   rR     s    zTrainer.from_argparse_args)
arg_parserr   c                 C   s
   t | |S r
  )rS   )r  r  r   r   r   rS     s    zTrainer.parse_argparserc                 C   s   t | S r
  )rT   )r  r   r   r   match_env_arguments  s    zTrainer.match_env_arguments)parent_parserrS  r   c                 K   s   t | |fi |S r
  )rQ   )r  r  rS  r   r   r   rQ     s    zTrainer.add_argparse_argsc                 C   s   | j jtjkS r
  )r   r   rJ   ZINTERRUPTEDr   r   r   r   interrupted  s    zTrainer.interruptedc                 C   s   | j jtjkS r
  )r   r%  rG   rv  r   r   r   r   r     s    zTrainer.training)valr   c                 C   s"   |rt j| j_n| jrd | j_d S r
  )rG   rv  r   r%  r   r   r  r   r   r   r     s    c                 C   s   | j jtjkS r
  )r   r%  rG   r   r   r   r   r   r     s    zTrainer.testingc                 C   s"   |rt j| j_n| jrd | j_d S r
  )rG   r   r   r%  r   r  r   r   r   r     s    c                 C   s   | j jtjkS r
  )r   r%  rG   r  r   r   r   r   r    s    zTrainer.predictingc                 C   s"   |rt j| j_n| jrd | j_d S r
  )rG   r  r   r%  r  r  r   r   r   r    s    c                 C   s   t d | jjtjkS )NzM`Trainer.tuning` has been deprecated in v1.8.0 and will be removed in v2.0.0.)r[   r   r%  rG   TUNINGr   r   r   r   tuning  s    zTrainer.tuningc                 C   s*   t d |rtj| j_n| jr&d | j_d S )NzUSetting `Trainer.tuning` has been deprecated in v1.8.0 and will be removed in v2.0.0.)r[   rG   r  r   r%  r  r  r   r   r   r    s
    c                 C   s   | j jtjkS r
  )r   r%  rG   r   r   r   r   r   r     s    zTrainer.validatingc                 C   s"   |rt j| j_n| jrd | j_d S r
  )rG   r   r   r%  r   r  r   r   r   r     s    c                 C   s   | j jd uo| j jjS r
  )r   r%  r3  r   r   r   r   r3    s    zTrainer.evaluatingc                 C   s   | j jtjkS r
  )r   r%  rG   SANITY_CHECKINGr   r   r   r   rH    s    zTrainer.sanity_checkingc                 C   s"   |rt j| j_n| jrd | j_d S r
  )rG   r  r   r%  rH  r  r   r   r   rH    s    c                 C   s
   | j jjS )zThe number of optimizer steps taken (does not reset each epoch).

        This includes multiple optimizers and TBPTT steps (if enabled).
        )r   r   global_stepr   r   r   r   r    s    zTrainer.global_stepc                 C   s   | j jjjS )z=The current epoch, updated after the epoch end hooks are run.)r   Zepoch_progresscurrentZ	completedr   r   r   r   rz    s    zTrainer.current_epochc                 C   s   | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.max_epochsc                 C   s   | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.min_epochsc                 C   s   | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.max_stepsc                 C   s   | j jS r
  )r   r   r   r   r   r   r     s    zTrainer.min_stepsc                 C   s   | j jjjS )z,Whether trainer is executing the last batch.)r   r   Zbatch_progressis_last_batchr   r   r   r   r     s    zTrainer.is_last_batchc                 C   s   | j S r
  )	_fit_loopr   r   r   r   r   %  s    zTrainer.fit_loop)r0  r   c                 C   s   | |_ || _dS )zAttach a custom fit loop to this Trainer.

        It will run with
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.fit`.
        N)r7  r  r/  r   r   r   r   )  s    c                 C   s   | j S r
  )_validate_loopr   r   r   r   r   3  s    zTrainer.validate_loopc                 C   s   | |_ || _dS )a-  Attach a custom validation loop to this Trainer.

        It will run with
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.validate`. Note that this loop is different from the one
        running during training inside the :meth:`pytorch_lightning.trainer.trainer.Trainer.fit` call.
        N)r7  r  r/  r   r   r   r   7  s    c                 C   s   | j S r
  )
_test_loopr   r   r   r   r   B  s    zTrainer.test_loopc                 C   s   | |_ || _dS )zAttach a custom test loop to this Trainer.

        It will run with
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.test`.
        N)r7  r  r/  r   r   r   r   F  s    c                 C   s   | j S r
  )_predict_loopr   r   r   r   r   P  s    zTrainer.predict_loopc                 C   s   | |_ || _dS )zAttach a custom prediction loop to this Trainer.

        It will run with
        :meth:`~pytorch_lightning.trainer.trainer.Trainer.predict`.
        N)r7  r  r/  r   r   r   r   T  s    c                 C   sL   | j jtjkr| jjjS | j jtjkr,| jS | j jtj	kr@| j
S tdd S )NzPThe `Trainer._evaluation_loop` property isn't defined. Accessed outside of scope)r   r   rH   r   r   r   rF  r   r   r   r   r  r   r   r   r   r9  ^  s    
zTrainer._evaluation_loopc                 C   s.   | j r| jS | js| jr| jS | jr*| jS d S r
  )r   r   rH  r3  r9  r  r   r   r   r   r   r.  h  s    zTrainer._active_loopc                 C   s   t | jdkr| jd S d S r  )ry  r&  r   r   r   r   rn   u  s    zTrainer.logger)rn   r   c                 C   s   |sg | _ n|g| _ d S r
  )r&  )r   rn   r   r   r   rn   y  s    c                 C   s   | j S r
  Z_loggersr   r   r   r   r&    s    zTrainer.loggers)r&  r   c                 C   s   |r|ng | _ d S r
  r  )r   r&  r   r   r   r&    s    c                 C   s   | j jS r
  )r   callback_metricsr   r   r   r   r    s    zTrainer.callback_metricsc                 C   s   | j jS r
  )r   logged_metricsr   r   r   r   r    s    zTrainer.logged_metricsc                 C   s   | j jS r
  )r   progress_bar_metricsr   r   r   r   r    s    zTrainer.progress_bar_metricsc                 C   s   | j }|d ur|jS d S r
  )r.  _results)r   Zactive_loopr   r   r   r    s    zTrainer._resultsc                 C   s   t  r|  sd S tdd S r  )rY   _should_terminate_gracefullyrW   r   r   r   r   _exit_gracefully_on_signal  s    z"Trainer._exit_gracefully_on_signalc                 C   s0   t jt| j| jjd}t| jj|dddkS )N)r  sum)Z	reduce_opr   )r6  Ztensorr   Z_terminate_gracefullyr   r  boolreduce)r   valuer   r   r   r    s    z$Trainer._should_terminate_gracefullyc                 C   s   | j }|jdgkrtd| jdkr<| jdkr6tdS | jS | jdu rVtd |   | j	}|tdkrn| jS | jdus|J |
| j| _| j}t|| t| jd }| jdkrt|| jn|}|S )a  
        Estimated stepping batches for the complete training inferred from DataLoaders, gradient
        accumulation factor and distributed setup.

        Examples::

            def configure_optimizers(self):
                optimizer = ...
                scheduler = torch.optim.lr_scheduler.OneCycleLR(
                    optimizer, max_lr=1e-3, total_steps=self.trainer.estimated_stepping_batches
                )
                return [optimizer], [scheduler]

        r   zkEstimated stepping batches cannot be computed with different `accumulate_grad_batches` at different epochs.ri   r   NzBLoading `train_dataloader` to estimate number of stepping batches.rg   )accumulation_schedulerZepochsrX   r   r   r   r   r\   r  r   Zget_accumulate_grad_batchesrz  r   mathceilr{  rD  )r   r  Ztotal_batchesZeffective_batch_sizeZmax_estimated_stepsr   r   r   estimated_stepping_batches  s&    

z"Trainer.estimated_stepping_batches)3TTNNNNrg   NNNNNNTrh   ri   rg   FNNNri   NNNNNNNrj   NNFrk   Trl   NNNNr   FTFFNNNFrm   T)NNNN)NNNN)NNNTN)NNNTN)NNNTN)NNNTN)NNNNN)NNNNN)NNNNNNr   )N)N)N)N)N)N)FN)r   
__module__r   rP   r   r&   r   r  r   r   r    r   r   r   strr
   r   r   r7   r<   r   r2   r;   r0   r   r   objectr   rc   r%   ra   r   r   r_   r   r`   r   r   r   r  r   r	   r   rL   r  r  r   r!  r"  r  r4  r  r  r  r5  r  r  r#  r   rK  r  r=  r]  ra  rl  rm  rn  staticmethodr   r  r  r  r  rC  propertyr   r   r1   r  rx  rp  r  rr  rt   r  r  r   r   r  setterrb   r  r  r   r=   r   r  r6  nnModuler   rq  r  r  r  rG  rq   r"   r  r  r$   r  r!   r  r  r#   r  r   r   r  classmethodr  r   r   rR   rS   r  r   rQ   r  r   r   r  r  r   r3  rH  r  rz  r   r   r   r   r  r+   r   r*   r   r   r(   r   r9  r.  rn   r&  r  rC   r  rD   r  rE   r  r  r  r  __classcell__r   r   r   r   rf   v   s                                                    



   ,   #   .    /    5    0    5    1    4      3'
-

"
s	

	
			rf   T)r   r   r   c                 c   sf   |r0t  r t  r t  dks0t| ts0tjntj}|  d V  W d    n1 sX0    Y  d S )NZgloo)	distZis_availableZis_initializedget_backendr   r   r6  r   rI  )r   r   Zcontext_manager_classr   r   r   r<    s    
r<  )T)__doc__r  loggingr  r   warningsargparser   r   r   
contextlibr   copyr   datetimer   pathlibr   typingr	   r
   r   r   r   r   r   r   weakrefr   r6  Ztorch.distributedZdistributedr  Z#lightning_utilities.core.apply_funcr   Z lightning_utilities.core.importsr   Zpackaging.versionr   r   Ztorch.optimr   Ztorch.utils.datar   Ztyping_extensionsr   Zpytorch_lightningr   Z#lightning_fabric.utilities.cloud_ior   Zlightning_fabric.utilities.datar   Z"lightning_fabric.utilities.importsr   Z lightning_fabric.utilities.typesr   Z#lightning_fabric.utilities.warningsr   Zpytorch_lightning.acceleratorsr   r   Zpytorch_lightning.callbacksr    r!   r"   r#   Z-pytorch_lightning.callbacks.prediction_writerr$   Z!pytorch_lightning.core.datamoduler%   Zpytorch_lightning.loggersr&   Z%pytorch_lightning.loggers.tensorboardr'   Zpytorch_lightning.loopsr(   r)   Z2pytorch_lightning.loops.dataloader.evaluation_loopr*   Z pytorch_lightning.loops.fit_loopr+   Z!pytorch_lightning.loops.utilitiesr,   r-   Zpytorch_lightning.pluginsr.   r/   r0   r1   Zpytorch_lightning.profilersr2   Zpytorch_lightning.strategiesr3   r4   r5   r6   r7   Zpytorch_lightning.trainerr8   r9   Z1pytorch_lightning.trainer.configuration_validatorr:   Z:pytorch_lightning.trainer.connectors.accelerator_connectorr;   r<   r=   r>   Z7pytorch_lightning.trainer.connectors.callback_connectorr?   Z9pytorch_lightning.trainer.connectors.checkpoint_connectorr@   Z3pytorch_lightning.trainer.connectors.data_connectorrA   Z5pytorch_lightning.trainer.connectors.logger_connectorrB   Z<pytorch_lightning.trainer.connectors.logger_connector.resultrC   rD   rE   Z5pytorch_lightning.trainer.connectors.signal_connectorrF   Z pytorch_lightning.trainer.statesrG   rH   rI   rJ   Z$pytorch_lightning.trainer.supportersrK   Zpytorch_lightning.tuner.tuningrL   rM   Zpytorch_lightning.utilitiesrN   rO   Z$pytorch_lightning.utilities.argparserP   rQ   rR   rS   rT   Z(pytorch_lightning.utilities.auto_restartrU   Z pytorch_lightning.utilities.datarV   Z&pytorch_lightning.utilities.exceptionsrW   rX   Z#pytorch_lightning.utilities.importsrY   Z)pytorch_lightning.utilities.model_helpersrZ   Z%pytorch_lightning.utilities.rank_zeror[   r\   r]   Z pytorch_lightning.utilities.seedr^   Z!pytorch_lightning.utilities.typesr_   r`   ra   rb   rc   	getLoggerr   r   filterwarningsrf   r  r<  r   r   r   r   <module>   s   (
                r