mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Docs11 (#1054)
* training_end to training_step_end * training_end to training_step_end * training_end to training_step_end * training_end to training_step_end * training_end to training_step_end
This commit is contained in:
@@ -98,26 +98,78 @@ Use it to do whatever!
|
|||||||
Trainer flags
|
Trainer flags
|
||||||
-------------
|
-------------
|
||||||
|
|
||||||
logger
|
accumulate_grad_batches
|
||||||
^^^^^^
|
^^^^^^^^^^^^^^^^^^^^^^^
|
||||||
|
Accumulates grads every k batches or as set up in the dict.
|
||||||
Logger (or iterable collection of loggers) for experiment tracking.
|
|
||||||
|
|
||||||
.. code-block:: python
|
|
||||||
|
|
||||||
Trainer(logger=logger)
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
from pytorch_lightning.loggers import TensorBoardLogger
|
# default used by the Trainer (no accumulation)
|
||||||
|
trainer = Trainer(accumulate_grad_batches=1)
|
||||||
|
|
||||||
# default logger used by trainer
|
# accumulate every 4 batches (effective batch size is batch*4)
|
||||||
logger = TensorBoardLogger(
|
trainer = Trainer(accumulate_grad_batches=4)
|
||||||
save_dir=os.getcwd(),
|
|
||||||
version=self.slurm_job_id,
|
|
||||||
name='lightning_logs'
|
|
||||||
)
|
|
||||||
|
|
||||||
|
# no accumulation for epochs 1-4. accumulate 3 for epochs 5-10. accumulate 20 after that
|
||||||
|
trainer = Trainer(accumulate_grad_batches={5: 3, 10: 20})
|
||||||
|
|
||||||
|
amp_level
|
||||||
|
^^^^^^^^^
|
||||||
|
The optimization level to use (O1, O2, etc...)
|
||||||
|
for 16-bit GPU precision (using NVIDIA apex under the hood).
|
||||||
|
|
||||||
|
Check nvidia docs for level (https://nvidia.github.io/apex/amp.html#opt-levels)
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(amp_level='O1')
|
||||||
|
|
||||||
|
benchmark
|
||||||
|
^^^^^^^^^
|
||||||
|
|
||||||
|
If true enables cudnn.benchmark.
|
||||||
|
This flag is likely to increase the speed of your system if your
|
||||||
|
input sizes don't change. However, if it does, then it will likely
|
||||||
|
make your system slower.
|
||||||
|
|
||||||
|
The speedup comes from allowing the cudnn auto-tuner to find the best
|
||||||
|
algorithm for the hardware `[see discussion here]
|
||||||
|
<https://discuss.pytorch.org/t/what-does-torch-backends-cudnn-benchmark-do/5936>`_.
|
||||||
|
|
||||||
|
callbacks
|
||||||
|
^^^^^^^^^
|
||||||
|
|
||||||
|
callbacks: Add a list of callbacks.
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
# a list of callbacks
|
||||||
|
callbacks = [PrintCallback()]
|
||||||
|
trainer = Trainer(callbacks=callbacks)
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
from pytorch_lightning.callbacks import Callback
|
||||||
|
|
||||||
|
class PrintCallback(Callback):
|
||||||
|
def on_train_start(self):
|
||||||
|
print("Training is started!")
|
||||||
|
def on_train_end(self):
|
||||||
|
print(f"Training is done. The logs are: {self.trainer.logs}")
|
||||||
|
|
||||||
|
check_val_every_n_epoch
|
||||||
|
^^^^^^^^^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
|
Check val every n train epochs.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(check_val_every_n_epoch=1)
|
||||||
|
|
||||||
|
# run val loop every 10 training epochs
|
||||||
|
trainer = Trainer(check_val_every_n_epoch=10)
|
||||||
|
|
||||||
checkpoint_callback
|
checkpoint_callback
|
||||||
^^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^^^
|
||||||
@@ -141,6 +193,43 @@ Example::
|
|||||||
prefix=''
|
prefix=''
|
||||||
)
|
)
|
||||||
|
|
||||||
|
default_save_path
|
||||||
|
^^^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
|
Default path for logs and weights when no logger/ckpt_callback passed
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(default_save_path=os.getcwd())
|
||||||
|
|
||||||
|
distributed_backend
|
||||||
|
^^^^^^^^^^^^^^^^^^^
|
||||||
|
The distributed backend to use.
|
||||||
|
|
||||||
|
- ('dp') is DataParallel (split batch among GPUs of same machine)
|
||||||
|
- ('ddp') is DistributedDataParallel (each gpu on each node trains, and syncs grads)
|
||||||
|
- ('ddp2') dp on node, ddp across nodes
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(distributed_backend=None)
|
||||||
|
|
||||||
|
# dp = DataParallel (split a batch onto k gpus on same machine).
|
||||||
|
trainer = Trainer(gpus=2, distributed_backend='dp')
|
||||||
|
|
||||||
|
# ddp = DistributedDataParallel
|
||||||
|
# Each gpu trains by itself on a subset of the data.
|
||||||
|
# Gradients sync across all gpus and all machines.
|
||||||
|
trainer = Trainer(gpus=2, num_nodes=2, distributed_backend='ddp')
|
||||||
|
|
||||||
|
# ddp2 = DistributedDataParallel + dp
|
||||||
|
# behaves like dp on every node
|
||||||
|
# syncs gradients across nodes like ddp
|
||||||
|
# useful for things like increasing the number of negative samples
|
||||||
|
trainer = Trainer(gpus=2, num_nodes=2, distributed_backend='ddp2')
|
||||||
|
|
||||||
early_stop_callback
|
early_stop_callback
|
||||||
^^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
@@ -171,78 +260,35 @@ Example::
|
|||||||
mode='min'
|
mode='min'
|
||||||
)
|
)
|
||||||
|
|
||||||
callbacks
|
fast_dev_run
|
||||||
^^^^^^^^^
|
^^^^^^^^^^^^
|
||||||
|
|
||||||
callbacks: Add a list of callbacks.
|
Runs 1 batch of train, test and val to find any bugs (ie: a sort of unit test).
|
||||||
|
|
||||||
|
Under the hood the pseudocode looks like this:
|
||||||
|
|
||||||
.. code-block:: python
|
.. code-block:: python
|
||||||
|
|
||||||
# a list of callbacks
|
# loading
|
||||||
callbacks = [PrintCallback()]
|
__init__()
|
||||||
trainer = Trainer(callbacks=callbacks)
|
prepare_data
|
||||||
|
|
||||||
Example::
|
# test training step
|
||||||
|
training_batch = next(train_dataloader)
|
||||||
|
training_step(training_batch)
|
||||||
|
|
||||||
from pytorch_lightning.callbacks import Callback
|
# test val step
|
||||||
|
val_batch = next(val_dataloader)
|
||||||
class PrintCallback(Callback):
|
out = validation_step(val_batch)
|
||||||
def on_train_start(self):
|
validation_epoch_end([out])
|
||||||
print("Training is started!")
|
|
||||||
def on_train_end(self):
|
|
||||||
print(f"Training is done. The logs are: {self.trainer.logs}")
|
|
||||||
|
|
||||||
default_save_path
|
|
||||||
^^^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
Default path for logs and weights when no logger/ckpt_callback passed
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(default_save_path=os.getcwd())
|
|
||||||
|
|
||||||
gradient_clip_val
|
|
||||||
^^^^^^^^^^^^^^^^^
|
|
||||||
Gradient clipping value
|
|
||||||
|
|
||||||
- 0 means don't clip.
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(gradient_clip_val=0.0)
|
trainer = Trainer(fast_dev_run=False)
|
||||||
|
|
||||||
gradient_clip
|
# runs 1 train, val, test batch and program ends
|
||||||
.. warning: .. deprecated:: 0.5.0
|
trainer = Trainer(fast_dev_run=True)
|
||||||
Use `gradient_clip_val` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
process_position
|
|
||||||
^^^^^^^^^^^^^^^^
|
|
||||||
orders the tqdm bar when running multiple models on same machine.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(process_position=0)
|
|
||||||
|
|
||||||
num_nodes
|
|
||||||
^^^^^^^^^
|
|
||||||
|
|
||||||
Number of GPU nodes for distributed training.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(num_nodes=1)
|
|
||||||
|
|
||||||
# to train on 8 nodes
|
|
||||||
trainer = Trainer(num_nodes=8)
|
|
||||||
|
|
||||||
nb_gpu_nodes
|
|
||||||
|
|
||||||
..warning:: .. deprecated:: 0.5.0
|
|
||||||
Use `num_nodes` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
gpus
|
gpus
|
||||||
^^^^
|
^^^^
|
||||||
@@ -270,6 +316,158 @@ Example::
|
|||||||
# combine with num_nodes to train on multiple GPUs across nodes
|
# combine with num_nodes to train on multiple GPUs across nodes
|
||||||
trainer = Trainer(gpus=2, num_nodes=4) # uses 8 gpus in total
|
trainer = Trainer(gpus=2, num_nodes=4) # uses 8 gpus in total
|
||||||
|
|
||||||
|
gradient_clip_val
|
||||||
|
^^^^^^^^^^^^^^^^^
|
||||||
|
Gradient clipping value
|
||||||
|
|
||||||
|
- 0 means don't clip.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(gradient_clip_val=0.0)
|
||||||
|
|
||||||
|
gradient_clip
|
||||||
|
.. warning: .. deprecated:: 0.5.0
|
||||||
|
Use `gradient_clip_val` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
|
log_gpu_memory
|
||||||
|
^^^^^^^^^^^^^^
|
||||||
|
Options:
|
||||||
|
|
||||||
|
- None
|
||||||
|
- 'min_max'
|
||||||
|
- 'all'
|
||||||
|
|
||||||
|
.. note:: Might slow performance because it uses the output of nvidia-smi.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(log_gpu_memory=None)
|
||||||
|
|
||||||
|
# log all the GPUs (on master node only)
|
||||||
|
trainer = Trainer(log_gpu_memory='all')
|
||||||
|
|
||||||
|
# log only the min and max memory on the master node
|
||||||
|
trainer = Trainer(log_gpu_memory='min_max')
|
||||||
|
|
||||||
|
log_save_interval
|
||||||
|
^^^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
|
Writes logs to disk this often
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(log_save_interval=100)
|
||||||
|
|
||||||
|
logger
|
||||||
|
^^^^^^
|
||||||
|
|
||||||
|
Logger (or iterable collection of loggers) for experiment tracking.
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
Trainer(logger=logger)
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
from pytorch_lightning.loggers import TensorBoardLogger
|
||||||
|
|
||||||
|
# default logger used by trainer
|
||||||
|
logger = TensorBoardLogger(
|
||||||
|
save_dir=os.getcwd(),
|
||||||
|
version=self.slurm_job_id,
|
||||||
|
name='lightning_logs'
|
||||||
|
)
|
||||||
|
|
||||||
|
max_epochs
|
||||||
|
^^^^^^^^^^
|
||||||
|
Stop training once this number of epochs is reached
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(max_epochs=1000)
|
||||||
|
|
||||||
|
max_nb_epochs
|
||||||
|
|
||||||
|
.. warning:: .. deprecated:: 0.5.0
|
||||||
|
Use `max_epochs` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
|
min_epochs
|
||||||
|
^^^^^^^^^^
|
||||||
|
Force training for at least these many epochs
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(min_epochs=1)
|
||||||
|
|
||||||
|
min_nb_epochs:
|
||||||
|
|
||||||
|
.. warning:: deprecated:: 0.5.0
|
||||||
|
Use `min_nb_epochs` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
|
max_steps
|
||||||
|
^^^^^^^^^
|
||||||
|
Stop training after this number of steps. Disabled by default (None).
|
||||||
|
Training will stop if max_steps or max_epochs have reached (earliest).
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# Stop after 100 steps
|
||||||
|
trainer = Trainer(max_steps=100)
|
||||||
|
|
||||||
|
min_steps
|
||||||
|
^^^^^^^^^
|
||||||
|
|
||||||
|
Force training for at least these number of steps. Disabled by default (None).
|
||||||
|
Trainer will train model for at least min_steps or min_epochs (latest).
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# Run at least for 100 steps (disable min_epochs)
|
||||||
|
trainer = Trainer(min_steps=100, min_epochs=0)
|
||||||
|
|
||||||
|
num_nodes
|
||||||
|
^^^^^^^^^
|
||||||
|
|
||||||
|
Number of GPU nodes for distributed training.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(num_nodes=1)
|
||||||
|
|
||||||
|
# to train on 8 nodes
|
||||||
|
trainer = Trainer(num_nodes=8)
|
||||||
|
|
||||||
|
nb_gpu_nodes
|
||||||
|
|
||||||
|
.. warning:: .. deprecated:: 0.5.0
|
||||||
|
Use `num_nodes` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
|
num_sanity_val_steps
|
||||||
|
^^^^^^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
|
Sanity check runs n batches of val before starting the training routine.
|
||||||
|
This catches any bugs in your validation without having to wait for the first validation check.
|
||||||
|
The Trainer uses 5 steps by default. Turn it off or modify it here.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(num_sanity_val_steps=5)
|
||||||
|
|
||||||
|
# turn it off
|
||||||
|
trainer = Trainer(num_sanity_val_steps=0)
|
||||||
|
|
||||||
|
nb_sanity_val_steps:
|
||||||
|
.. warning:: .. deprecated:: 0.5.0
|
||||||
|
Use `num_sanity_val_steps` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
num_tpu_cores
|
num_tpu_cores
|
||||||
^^^^^^^^^^^^^
|
^^^^^^^^^^^^^
|
||||||
How many TPU cores to train on (1 or 8).
|
How many TPU cores to train on (1 or 8).
|
||||||
@@ -315,42 +513,6 @@ Example::
|
|||||||
--env=XLA_USE_BF16=1
|
--env=XLA_USE_BF16=1
|
||||||
-- python your_trainer_file.py
|
-- python your_trainer_file.py
|
||||||
|
|
||||||
log_gpu_memory
|
|
||||||
^^^^^^^^^^^^^^
|
|
||||||
Options:
|
|
||||||
|
|
||||||
- None
|
|
||||||
- 'min_max'
|
|
||||||
- 'all'
|
|
||||||
|
|
||||||
.. note:: Might slow performance because it uses the output of nvidia-smi.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(log_gpu_memory=None)
|
|
||||||
|
|
||||||
# log all the GPUs (on master node only)
|
|
||||||
trainer = Trainer(log_gpu_memory='all')
|
|
||||||
|
|
||||||
# log only the min and max memory on the master node
|
|
||||||
trainer = Trainer(log_gpu_memory='min_max')
|
|
||||||
|
|
||||||
show_progress_bar
|
|
||||||
^^^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
If true shows tqdm progress bar
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(show_progress_bar=True)
|
|
||||||
|
|
||||||
progress_bar_refresh_rate
|
|
||||||
^^^^^^^^^^^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
How often to refresh progress bar (in steps)
|
|
||||||
|
|
||||||
overfit_pct
|
overfit_pct
|
||||||
^^^^^^^^^^^
|
^^^^^^^^^^^
|
||||||
uses this much data of all datasets.
|
uses this much data of all datasets.
|
||||||
@@ -363,153 +525,127 @@ Example::
|
|||||||
# use only 1% of the train, test, val datasets
|
# use only 1% of the train, test, val datasets
|
||||||
trainer = Trainer(overfit_pct=0.01)
|
trainer = Trainer(overfit_pct=0.01)
|
||||||
|
|
||||||
track_grad_norm
|
precision
|
||||||
|
^^^^^^^^^
|
||||||
|
Full precision (32), half precision (16).
|
||||||
|
Can be used on CPU, GPU or TPUs.
|
||||||
|
|
||||||
|
If used on TPU will use torch.bfloat16 but tensor printing
|
||||||
|
will still show torch.float32.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(precision=32)
|
||||||
|
|
||||||
|
# 16-bit precision
|
||||||
|
trainer = Trainer(precision=16)
|
||||||
|
|
||||||
|
# one day
|
||||||
|
trainer = Trainer(precision=8|4|2)
|
||||||
|
|
||||||
|
print_nan_grads
|
||||||
^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
- no tracking (-1)
|
Prints gradients with nan values
|
||||||
- Otherwise tracks that norm (2 for 2-norm)
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(track_grad_norm=-1)
|
trainer = Trainer(print_nan_grads=False)
|
||||||
|
|
||||||
# track the 2-norm
|
process_position
|
||||||
trainer = Trainer(track_grad_norm=2)
|
^^^^^^^^^^^^^^^^
|
||||||
|
orders the tqdm bar when running multiple models on same machine.
|
||||||
check_val_every_n_epoch
|
|
||||||
^^^^^^^^^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
Check val every n train epochs.
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(check_val_every_n_epoch=1)
|
trainer = Trainer(process_position=0)
|
||||||
|
|
||||||
# run val loop every 10 training epochs
|
profiler
|
||||||
trainer = Trainer(check_val_every_n_epoch=10)
|
^^^^^^^^
|
||||||
|
To profile individual steps during training and assist in identifying bottlenecks.
|
||||||
|
|
||||||
fast_dev_run
|
Example::
|
||||||
^^^^^^^^^^^^
|
|
||||||
|
|
||||||
Runs 1 batch of train, test and val to find any bugs (ie: a sort of unit test).
|
from pytorch_lightning.profiler import Profiler, AdvancedProfiler
|
||||||
|
|
||||||
Under the hood the pseudocode looks like this:
|
# default used by the Trainer
|
||||||
|
trainer = Trainer(profiler=None)
|
||||||
|
|
||||||
|
# to profile standard training events
|
||||||
|
trainer = Trainer(profiler=True)
|
||||||
|
|
||||||
|
# equivalent to profiler=True
|
||||||
|
profiler = Profiler()
|
||||||
|
trainer = Trainer(profiler=profiler)
|
||||||
|
|
||||||
|
# advanced profiler for function-level stats
|
||||||
|
profiler = AdvancedProfiler()
|
||||||
|
trainer = Trainer(profiler=profiler)
|
||||||
|
|
||||||
|
progress_bar_refresh_rate
|
||||||
|
^^^^^^^^^^^^^^^^^^^^^^^^^
|
||||||
|
How often to refresh progress bar (in steps)
|
||||||
|
Default is 50. Useful for notebooks with slow refresh rate.
|
||||||
|
|
||||||
|
reload_dataloaders_every_epoch
|
||||||
|
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
|
||||||
|
Set to True to reload dataloaders every epoch
|
||||||
|
|
||||||
.. code-block:: python
|
.. code-block:: python
|
||||||
|
|
||||||
# loading
|
# if False (default)
|
||||||
__init__()
|
train_loader = model.train_dataloader()
|
||||||
prepare_data
|
for epoch in epochs:
|
||||||
|
for batch in train_loader:
|
||||||
|
...
|
||||||
|
|
||||||
# test training step
|
# if True
|
||||||
training_batch = next(train_dataloader)
|
for epoch in epochs:
|
||||||
training_step(training_batch)
|
train_loader = model.train_dataloader()
|
||||||
|
for batch in train_loader:
|
||||||
|
|
||||||
# test val step
|
resume_from_checkpoint
|
||||||
val_batch = next(val_dataloader)
|
^^^^^^^^^^^^^^^^^^^^^^
|
||||||
out = validation_step(val_batch)
|
To resume training from a specific checkpoint pass in the path here.k
|
||||||
validation_epoch_end([out])
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(fast_dev_run=False)
|
trainer = Trainer(resume_from_checkpoint=None)
|
||||||
|
|
||||||
# runs 1 train, val, test batch and program ends
|
# resume from a specific checkpoint
|
||||||
trainer = Trainer(fast_dev_run=True)
|
trainer = Trainer(resume_from_checkpoint='some/path/to/my_checkpoint.ckpt')
|
||||||
|
|
||||||
accumulate_grad_batches
|
row_log_interval
|
||||||
^^^^^^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^
|
||||||
Accumulates grads every k batches or as set up in the dict.
|
|
||||||
|
|
||||||
Example::
|
How often to add logging rows (does not write to disk)
|
||||||
|
|
||||||
# default used by the Trainer (no accumulation)
|
|
||||||
trainer = Trainer(accumulate_grad_batches=1)
|
|
||||||
|
|
||||||
# accumulate every 4 batches (effective batch size is batch*4)
|
|
||||||
trainer = Trainer(accumulate_grad_batches=4)
|
|
||||||
|
|
||||||
# no accumulation for epochs 1-4. accumulate 3 for epochs 5-10. accumulate 20 after that
|
|
||||||
trainer = Trainer(accumulate_grad_batches={5: 3, 10: 20})
|
|
||||||
|
|
||||||
max_epochs
|
|
||||||
^^^^^^^^^^
|
|
||||||
Stop training once this number of epochs is reached
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(max_epochs=1000)
|
trainer = Trainer(row_log_interval=10)
|
||||||
|
|
||||||
max_nb_epochs
|
|
||||||
|
|
||||||
|
add_row_log_interval
|
||||||
.. warning:: .. deprecated:: 0.5.0
|
.. warning:: .. deprecated:: 0.5.0
|
||||||
Use `max_epochs` instead. Will remove 0.8.0.
|
Use `row_log_interval` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
min_epochs
|
use_amp:
|
||||||
^^^^^^^^^^
|
.. warning:: .. deprecated:: 0.6.1
|
||||||
Force training for at least these many epochs
|
Use `precision` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
Example::
|
show_progress_bar
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(min_epochs=1)
|
|
||||||
|
|
||||||
min_nb_epochs:
|
|
||||||
.. warning:: .. deprecated:: 0.5.0
|
|
||||||
Use `min_nb_epochs` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
max_steps
|
|
||||||
^^^^^^^^^
|
|
||||||
Stop training after this number of steps. Disabled by default (None).
|
|
||||||
Training will stop if max_steps or max_epochs have reached (earliest).
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# Stop after 100 steps
|
|
||||||
trainer = Trainer(max_steps=100)
|
|
||||||
|
|
||||||
min_steps
|
|
||||||
^^^^^^^^^
|
|
||||||
|
|
||||||
Force training for at least these number of steps. Disabled by default (None).
|
|
||||||
Trainer will train model for at least min_steps or min_epochs (latest).
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# Run at least for 100 steps (disable min_epochs)
|
|
||||||
trainer = Trainer(min_steps=100, min_epochs=0)
|
|
||||||
|
|
||||||
train_percent_check
|
|
||||||
^^^^^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
How much of training dataset to check.
|
|
||||||
Useful when debugging or testing something that happens at the end of an epoch.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(train_percent_check=1.0)
|
|
||||||
|
|
||||||
# run through only 25% of the training set each epoch
|
|
||||||
trainer = Trainer(train_percent_check=0.25)
|
|
||||||
|
|
||||||
val_percent_check
|
|
||||||
^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
How much of validation dataset to check.
|
If true shows tqdm progress bar
|
||||||
Useful when debugging or testing something that happens at the end of an epoch.
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(val_percent_check=1.0)
|
trainer = Trainer(show_progress_bar=True)
|
||||||
|
|
||||||
# run through only 25% of the validation set each epoch
|
|
||||||
trainer = Trainer(val_percent_check=0.25)
|
|
||||||
|
|
||||||
test_percent_check
|
test_percent_check
|
||||||
^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^^
|
||||||
@@ -545,156 +681,33 @@ Example::
|
|||||||
# (ie: production cases with streaming data)
|
# (ie: production cases with streaming data)
|
||||||
trainer = Trainer(val_check_interval=1000)
|
trainer = Trainer(val_check_interval=1000)
|
||||||
|
|
||||||
log_save_interval
|
track_grad_norm
|
||||||
^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^
|
||||||
|
|
||||||
Writes logs to disk this often
|
- no tracking (-1)
|
||||||
|
- Otherwise tracks that norm (2 for 2-norm)
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(log_save_interval=100)
|
trainer = Trainer(track_grad_norm=-1)
|
||||||
|
|
||||||
row_log_interval
|
# track the 2-norm
|
||||||
^^^^^^^^^^^^^^^^
|
trainer = Trainer(track_grad_norm=2)
|
||||||
|
|
||||||
How often to add logging rows (does not write to disk)
|
train_percent_check
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(row_log_interval=10)
|
|
||||||
|
|
||||||
add_row_log_interval
|
|
||||||
.. warning:: .. deprecated:: 0.5.0
|
|
||||||
Use `row_log_interval` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
distributed_backend
|
|
||||||
^^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^^^
|
||||||
The distributed backend to use.
|
|
||||||
|
|
||||||
- ('dp') is DataParallel (split batch among GPUs of same machine)
|
How much of training dataset to check.
|
||||||
- ('ddp') is DistributedDataParallel (each gpu on each node trains, and syncs grads)
|
Useful when debugging or testing something that happens at the end of an epoch.
|
||||||
- ('ddp2') dp on node, ddp across nodes
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(distributed_backend=None)
|
trainer = Trainer(train_percent_check=1.0)
|
||||||
|
|
||||||
# dp = DataParallel (split a batch onto k gpus on same machine).
|
# run through only 25% of the training set each epoch
|
||||||
trainer = Trainer(gpus=2, distributed_backend='dp')
|
trainer = Trainer(train_percent_check=0.25)
|
||||||
|
|
||||||
# ddp = DistributedDataParallel
|
|
||||||
# Each gpu trains by itself on a subset of the data.
|
|
||||||
# Gradients sync across all gpus and all machines.
|
|
||||||
trainer = Trainer(gpus=2, num_nodes=2, distributed_backend='ddp')
|
|
||||||
|
|
||||||
# ddp2 = DistributedDataParallel + dp
|
|
||||||
# behaves like dp on every node
|
|
||||||
# syncs gradients across nodes like ddp
|
|
||||||
# useful for things like increasing the number of negative samples
|
|
||||||
trainer = Trainer(gpus=2, num_nodes=2, distributed_backend='ddp2')
|
|
||||||
|
|
||||||
use_amp:
|
|
||||||
.. warning:: .. deprecated:: 0.6.1
|
|
||||||
Use `precision` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
precision
|
|
||||||
^^^^^^^^^
|
|
||||||
Full precision (32), half precision (16).
|
|
||||||
Can be used on CPU, GPU or TPUs.
|
|
||||||
|
|
||||||
If used on TPU will use torch.bfloat16 but tensor printing
|
|
||||||
will still show torch.float32.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(precision=32)
|
|
||||||
|
|
||||||
# 16-bit precision
|
|
||||||
trainer = Trainer(precision=16)
|
|
||||||
|
|
||||||
# one day
|
|
||||||
trainer = Trainer(precision=8|4|2)
|
|
||||||
|
|
||||||
print_nan_grads
|
|
||||||
^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
Prints gradients with nan values
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(print_nan_grads=False)
|
|
||||||
|
|
||||||
weights_summary
|
|
||||||
^^^^^^^^^^^^^^^
|
|
||||||
Prints a summary of the weights when training begins.
|
|
||||||
Options: 'full', 'top', None.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer (ie: print all weights)
|
|
||||||
trainer = Trainer(weights_summary='full')
|
|
||||||
|
|
||||||
# print only the top level modules
|
|
||||||
trainer = Trainer(weights_summary='top')
|
|
||||||
|
|
||||||
# don't print a summary
|
|
||||||
trainer = Trainer(weights_summary=None)
|
|
||||||
|
|
||||||
weights_save_path
|
|
||||||
^^^^^^^^^^^^^^^^^
|
|
||||||
Where to save weights if specified.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(weights_save_path=os.getcwd())
|
|
||||||
|
|
||||||
# save to your custom path
|
|
||||||
trainer = Trainer(weights_save_path='my/path')
|
|
||||||
|
|
||||||
# if checkpoint callback used, then overrides the weights path
|
|
||||||
# **NOTE: this saves weights to some/path NOT my/path
|
|
||||||
checkpoint_callback = ModelCheckpoint(filepath='some/path')
|
|
||||||
trainer = Trainer(
|
|
||||||
checkpoint_callback=checkpoint_callback,
|
|
||||||
weights_save_path='my/path'
|
|
||||||
)
|
|
||||||
|
|
||||||
amp_level
|
|
||||||
^^^^^^^^^
|
|
||||||
The optimization level to use (O1, O2, etc...)
|
|
||||||
for 16-bit GPU precision (using NVIDIA apex under the hood).
|
|
||||||
|
|
||||||
Check nvidia docs for level (https://nvidia.github.io/apex/amp.html#opt-levels)
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(amp_level='O1')
|
|
||||||
|
|
||||||
num_sanity_val_steps
|
|
||||||
^^^^^^^^^^^^^^^^^^^^
|
|
||||||
|
|
||||||
Sanity check runs n batches of val before starting the training routine.
|
|
||||||
This catches any bugs in your validation without having to wait for the first validation check.
|
|
||||||
The Trainer uses 5 steps by default. Turn it off or modify it here.
|
|
||||||
|
|
||||||
Example::
|
|
||||||
|
|
||||||
# default used by the Trainer
|
|
||||||
trainer = Trainer(num_sanity_val_steps=5)
|
|
||||||
|
|
||||||
# turn it off
|
|
||||||
trainer = Trainer(num_sanity_val_steps=0)
|
|
||||||
|
|
||||||
nb_sanity_val_steps:
|
|
||||||
.. warning:: .. deprecated:: 0.5.0
|
|
||||||
Use `num_sanity_val_steps` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
truncated_bptt_steps
|
truncated_bptt_steps
|
||||||
^^^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^^^^
|
||||||
@@ -723,69 +736,58 @@ Lightning takes care to split your batch along the time-dimension.
|
|||||||
.. note:: Using this feature requires updating your LightningModule's
|
.. note:: Using this feature requires updating your LightningModule's
|
||||||
:meth:`pytorch_lightning.core.LightningModule.training_step` to include a `hiddens` arg.
|
:meth:`pytorch_lightning.core.LightningModule.training_step` to include a `hiddens` arg.
|
||||||
|
|
||||||
resume_from_checkpoint
|
val_percent_check
|
||||||
^^^^^^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^
|
||||||
To resume training from a specific checkpoint pass in the path here.k
|
|
||||||
|
How much of validation dataset to check.
|
||||||
|
Useful when debugging or testing something that happens at the end of an epoch.
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(resume_from_checkpoint=None)
|
trainer = Trainer(val_percent_check=1.0)
|
||||||
|
|
||||||
# resume from a specific checkpoint
|
# run through only 25% of the validation set each epoch
|
||||||
trainer = Trainer(resume_from_checkpoint='some/path/to/my_checkpoint.ckpt')
|
trainer = Trainer(val_percent_check=0.25)
|
||||||
|
|
||||||
profiler
|
weights_save_path
|
||||||
^^^^^^^^
|
^^^^^^^^^^^^^^^^^
|
||||||
To profile individual steps during training and assist in identifying bottlenecks.
|
Where to save weights if specified.
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
from pytorch_lightning.profiler import Profiler, AdvancedProfiler
|
|
||||||
|
|
||||||
# default used by the Trainer
|
# default used by the Trainer
|
||||||
trainer = Trainer(profiler=None)
|
trainer = Trainer(weights_save_path=os.getcwd())
|
||||||
|
|
||||||
# to profile standard training events
|
# save to your custom path
|
||||||
trainer = Trainer(profiler=True)
|
trainer = Trainer(weights_save_path='my/path')
|
||||||
|
|
||||||
# equivalent to profiler=True
|
# if checkpoint callback used, then overrides the weights path
|
||||||
profiler = Profiler()
|
# **NOTE: this saves weights to some/path NOT my/path
|
||||||
trainer = Trainer(profiler=profiler)
|
checkpoint_callback = ModelCheckpoint(filepath='some/path')
|
||||||
|
trainer = Trainer(
|
||||||
|
checkpoint_callback=checkpoint_callback,
|
||||||
|
weights_save_path='my/path'
|
||||||
|
)
|
||||||
|
|
||||||
# advanced profiler for function-level stats
|
weights_summary
|
||||||
profiler = AdvancedProfiler()
|
^^^^^^^^^^^^^^^
|
||||||
trainer = Trainer(profiler=profiler)
|
Prints a summary of the weights when training begins.
|
||||||
|
Options: 'full', 'top', None.
|
||||||
|
|
||||||
|
Example::
|
||||||
|
|
||||||
reload_dataloaders_every_epoch
|
# default used by the Trainer (ie: print all weights)
|
||||||
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
|
trainer = Trainer(weights_summary='full')
|
||||||
Set to True to reload dataloaders every epoch
|
|
||||||
|
|
||||||
.. code-block:: python
|
# print only the top level modules
|
||||||
|
trainer = Trainer(weights_summary='top')
|
||||||
|
|
||||||
# if False (default)
|
# don't print a summary
|
||||||
train_loader = model.train_dataloader()
|
trainer = Trainer(weights_summary=None)
|
||||||
for epoch in epochs:
|
|
||||||
for batch in train_loader:
|
|
||||||
...
|
|
||||||
|
|
||||||
# if True
|
Trainer class
|
||||||
for epoch in epochs:
|
-------------
|
||||||
train_loader = model.train_dataloader()
|
|
||||||
for batch in train_loader:
|
|
||||||
|
|
||||||
benchmark
|
|
||||||
^^^^^^^^^
|
|
||||||
|
|
||||||
If true enables cudnn.benchmark.
|
|
||||||
This flag is likely to increase the speed of your system if your
|
|
||||||
input sizes don't change. However, if it does, then it will likely
|
|
||||||
make your system slower.
|
|
||||||
|
|
||||||
The speedup comes from allowing the cudnn auto-tuner to find the best
|
|
||||||
algorithm for the hardware `[see discussion here]
|
|
||||||
<https://discuss.pytorch.org/t/what-does-torch-backends-cudnn-benchmark-do/5936>`_.
|
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|||||||
@@ -126,71 +126,108 @@ class Trainer(TrainerIOMixin,
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
logger: Logger (or iterable collection of loggers) for experiment tracking.
|
logger: Logger (or iterable collection of loggers) for experiment tracking.
|
||||||
|
|
||||||
checkpoint_callback: Callback for checkpointing.
|
checkpoint_callback: Callback for checkpointing.
|
||||||
|
|
||||||
early_stop_callback (:class:`pytorch_lightning.callbacks.EarlyStopping`):
|
early_stop_callback (:class:`pytorch_lightning.callbacks.EarlyStopping`):
|
||||||
|
|
||||||
callbacks: Add a list of callbacks.
|
callbacks: Add a list of callbacks.
|
||||||
|
|
||||||
default_save_path: Default path for logs and weights when no logger/ckpt_callback passed
|
default_save_path: Default path for logs and weights when no logger/ckpt_callback passed
|
||||||
|
|
||||||
gradient_clip_val: 0 means don't clip.
|
gradient_clip_val: 0 means don't clip.
|
||||||
|
|
||||||
gradient_clip:
|
gradient_clip:
|
||||||
.. warning:: .. deprecated:: 0.6.1
|
.. warning:: deprecated 0.6.1 Use `gradient_clip_val` instead. Will remove 0.8.0.
|
||||||
Use `gradient_clip_val` instead. Will remove 0.8.0.
|
|
||||||
|
|
||||||
process_position: orders the tqdm bar when running multiple models on same machine.
|
process_position: orders the tqdm bar when running multiple models on same machine.
|
||||||
|
|
||||||
num_nodes: number of GPU nodes for distributed training.
|
num_nodes: number of GPU nodes for distributed training.
|
||||||
|
|
||||||
nb_gpu_nodes:
|
nb_gpu_nodes:
|
||||||
.. warning:: .. deprecated:: 0.6.1
|
.. warning:: .. deprecated:: 0.6.1
|
||||||
Use `num_nodes` instead. Will remove 0.8.0.
|
Use `num_nodes` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
gpus: Which GPUs to train on.
|
gpus: Which GPUs to train on.
|
||||||
|
|
||||||
num_tpu_cores: How many TPU cores to train on (1 or 8).
|
num_tpu_cores: How many TPU cores to train on (1 or 8).
|
||||||
|
|
||||||
log_gpu_memory: None, 'min_max', 'all'. Might slow performance
|
log_gpu_memory: None, 'min_max', 'all'. Might slow performance
|
||||||
|
|
||||||
show_progress_bar: If true shows tqdm progress bar
|
show_progress_bar: If true shows tqdm progress bar
|
||||||
|
|
||||||
progress_bar_refresh_rate: How often to refresh progress bar (in steps)
|
progress_bar_refresh_rate: How often to refresh progress bar (in steps)
|
||||||
|
|
||||||
track_grad_norm: -1 no tracking. Otherwise tracks that norm
|
track_grad_norm: -1 no tracking. Otherwise tracks that norm
|
||||||
|
|
||||||
check_val_every_n_epoch: Check val every n train epochs.
|
check_val_every_n_epoch: Check val every n train epochs.
|
||||||
|
|
||||||
fast_dev_run: runs 1 batch of train, test and val to find any bugs (ie: a sort of unit test).
|
fast_dev_run: runs 1 batch of train, test and val to find any bugs (ie: a sort of unit test).
|
||||||
|
|
||||||
accumulate_grad_batches: Accumulates grads every k batches or as set up in the dict.
|
accumulate_grad_batches: Accumulates grads every k batches or as set up in the dict.
|
||||||
|
|
||||||
max_epochs: Stop training once this number of epochs is reached.
|
max_epochs: Stop training once this number of epochs is reached.
|
||||||
|
|
||||||
max_nb_epochs:
|
max_nb_epochs:
|
||||||
.. warning:: .. deprecated:: 0.6.1
|
.. warning:: .. deprecated:: 0.6.1
|
||||||
Use `max_epochs` instead. Will remove 0.8.0.
|
Use `max_epochs` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
min_epochs: Force training for at least these many epochs
|
min_epochs: Force training for at least these many epochs
|
||||||
|
|
||||||
min_nb_epochs:
|
min_nb_epochs:
|
||||||
.. warning:: .. deprecated:: 0.6.1
|
.. warning:: .. deprecated:: 0.6.1
|
||||||
Use `min_epochs` instead. Will remove 0.8.0.
|
Use `min_epochs` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
max_steps: Stop training after this number of steps. Disabled by default (None).
|
max_steps: Stop training after this number of steps. Disabled by default (None).
|
||||||
|
|
||||||
min_steps: Force training for at least these number of steps. Disabled by default (None).
|
min_steps: Force training for at least these number of steps. Disabled by default (None).
|
||||||
|
|
||||||
train_percent_check: How much of training dataset to check.
|
train_percent_check: How much of training dataset to check.
|
||||||
|
|
||||||
val_percent_check: How much of validation dataset to check.
|
val_percent_check: How much of validation dataset to check.
|
||||||
|
|
||||||
test_percent_check: How much of test dataset to check.
|
test_percent_check: How much of test dataset to check.
|
||||||
|
|
||||||
val_check_interval: How often within one training epoch to check the validation set
|
val_check_interval: How often within one training epoch to check the validation set
|
||||||
|
|
||||||
log_save_interval: Writes logs to disk this often
|
log_save_interval: Writes logs to disk this often
|
||||||
|
|
||||||
row_log_interval: How often to add logging rows (does not write to disk)
|
row_log_interval: How often to add logging rows (does not write to disk)
|
||||||
|
|
||||||
add_row_log_interval:
|
add_row_log_interval:
|
||||||
.. warning:: .. deprecated:: 0.6.1
|
.. warning:: .. deprecated:: 0.6.1
|
||||||
Use `row_log_interval` instead. Will remove 0.8.0.
|
Use `row_log_interval` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
distributed_backend: The distributed backend to use.
|
distributed_backend: The distributed backend to use.
|
||||||
|
|
||||||
use_amp:
|
use_amp:
|
||||||
.. warning:: .. deprecated:: 0.7.0
|
.. warning:: .. deprecated:: 0.7.0
|
||||||
Use `precision` instead. Will remove 0.8.0.
|
Use `precision` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
precision: Full precision (32), half precision (16).
|
precision: Full precision (32), half precision (16).
|
||||||
|
|
||||||
print_nan_grads: Prints gradients with nan values
|
print_nan_grads: Prints gradients with nan values
|
||||||
|
|
||||||
weights_summary: Prints a summary of the weights when training begins.
|
weights_summary: Prints a summary of the weights when training begins.
|
||||||
|
|
||||||
weights_save_path: Where to save weights if specified.
|
weights_save_path: Where to save weights if specified.
|
||||||
|
|
||||||
amp_level: The optimization level to use (O1, O2, etc...).
|
amp_level: The optimization level to use (O1, O2, etc...).
|
||||||
|
|
||||||
num_sanity_val_steps: Sanity check runs n batches of val before starting the training routine.
|
num_sanity_val_steps: Sanity check runs n batches of val before starting the training routine.
|
||||||
|
|
||||||
nb_sanity_val_steps:
|
nb_sanity_val_steps:
|
||||||
.. warning:: .. deprecated:: 0.7.0
|
.. warning:: .. deprecated:: 0.7.0
|
||||||
Use `num_sanity_val_steps` instead. Will remove 0.8.0.
|
Use `num_sanity_val_steps` instead. Will remove 0.8.0.
|
||||||
|
|
||||||
truncated_bptt_steps: Truncated back prop breaks performs backprop every k steps of
|
truncated_bptt_steps: Truncated back prop breaks performs backprop every k steps of
|
||||||
|
|
||||||
resume_from_checkpoint: To resume training from a specific checkpoint pass in the path here.k
|
resume_from_checkpoint: To resume training from a specific checkpoint pass in the path here.k
|
||||||
|
|
||||||
profiler: To profile individual steps during training and assist in
|
profiler: To profile individual steps during training and assist in
|
||||||
|
|
||||||
reload_dataloaders_every_epoch: Set to True to reload dataloaders every epoch
|
reload_dataloaders_every_epoch: Set to True to reload dataloaders every epoch
|
||||||
|
|
||||||
benchmark (bool): If true enables cudnn.benchmark.
|
benchmark (bool): If true enables cudnn.benchmark.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user