Builds a Torch Learner from a ModelDescriptor and trains it with the given parameter specification.
The task type must be specified during construction.
Parameters
General:
The parameters of the optimizer, loss and callbacks,
prefixed with "opt.", "loss." and "cb.<callback id>." respectively, as well as:
epochs::integer(1)
The number of epochs.device::character(1)
The device. One of"auto","cpu", or"cuda"or other values defined inmlr_reflections$torch$devices. The value is initialized to"auto", which will select"cuda"if possible, then try"mps"and otherwise fall back to"cpu".num_threads::integer(1)
The number of threads for intraop parallelization (ifdeviceis"cpu"). This value is initialized to 1. When resampling, benchmarking or tuning in parallel, each worker uses this many threads, so divide the available cores among the workers instead of setting this to the number of cores.num_interop_threads::integer(1)
The number of threads for interop parallelization (ifdeviceis"cpu"). Note that this can only be set once per session, so setting this for one learner also changes the behavior of other learners, and a later learner asking for a different value errors.NULL(default) uses whatever is set. In order to use different values for this parameter, use encapsulation to train the learners in separate R sessions.seed::integer(1)or"random"orNULL
The torch seed that is used during training and prediction. This value is initialized to"random", which means that a random seed will be sampled at the beginning of the training phase. This seed (either set or randomly sampled) is available via$model$seedafter training and used during prediction. Note that by setting the seed during the training phase this will mean that by default (i.e. whenseedis"random"), clones of the learner will use a different seed. If set toNULL, no seeding will be done. This parameter only seeds torch's random number generator, it does not seed R's. Anything that is drawn from R's RNG is therefore unaffected by it, so to make those parts reproducible you need to seed R's RNG as well, e.g. viaset.seed().tensor_dataset::logical(1)|"device"
Whether to load all batches at once at the beginning of training and stack them. This is initialized toFALSE. If set to"device", the device of the tensors will be set to the value ofdevice, which can avoid unnecessary moving of tensors between devices. When your dataset fits into memory this will make the loading of batches faster. Note that this should not be set for datasets that containlazy_tensors with random data augmentation, as this augmentation will only be applied once at the beginning of training.
Evaluation:
measures_train::Measureorlist()ofMeasures
Measures to be evaluated during training.measures_valid::Measureorlist()ofMeasures
Measures to be evaluated during validation.eval_freq::integer(1)
How often the train / validation predictions are evaluated usingmeasures_train/measures_valid. This is initialized to1. Note that the final model is always evaluated.
Early Stopping:
patience::integer(1)
This activates early stopping using the validation scores. If the performance of a model does not improve forpatienceevaluation steps, training is ended. Note that this counts evaluation steps, not epochs: wheneval_freqis greater than1,patienceevaluation steps correspond topatience * eval_freqepochs. Note that the final model is stored in the learner, not the best model. This is initialized to0, which means no early stopping. The first entry frommeasures_validis used as the metric. This also requires to specify the$validatefield of the Learner, as well asmeasures_valid. If this is set, the epoch after which no improvement was observed, can be accessed via the$internal_tuned_valuesfield of the learner.min_delta::double(1)
The minimum improvement threshold for early stopping. Is initialized to 0.restore_best_weights::logical(1)
Whether to restore the weights of the best epoch when training ends, instead of keeping those of the last epoch that was trained. Is initialized toFALSE, i.e. the network of the last epoch is stored. Setting this toTRUEmakes the stored network the one of the epoch that$internal_tuned_valuesreports, and costs one additional copy of the network's parameters in memory. Checkpoints written byt_clbk("checkpoint")are unaffected: they always hold the network as training left it.
Dataloader:
batch_size::integer(1)
The batch size used by the training and prediction dataloader. It is required for training (unless abatch_sampleris provided, which already determines the batches) and it is required for prediction (unlessbatch_size_predictis set).batch_size_predict::integer(1)
The batch size used by the prediction dataloader (this includes the validation data during training). When set, it overridesbatch_sizefor prediction. The batch size does not change the predictions, but smaller batches take longer and require less memory.shuffle::logical(1)
Whether to shuffle the instances in the dataset. This is initialized toTRUE, which differs from the default (FALSE). It is ignored when asamplerorbatch_sampleris provided.sampler::torch::sampler
Object that defines how the dataloader draws samples, i.e. the order in which the observations are drawn. This must be the sampler generator (as returned bytorch::sampler()), not an instance, as it is instantiated with the training dataset internally.batch_sampler::torch::sampler
Object that defines how the dataloader draws batches. As forsampler, this must be the generator. When it is provided, the parametersbatch_size,shuffleanddrop_lastare ignored during training, because the batch sampler already determines the batches.num_workers::integer(1)
The number of workers for data loading (batches are loaded in parallel). The default is0, which means that data will be loaded in the main process.collate_fn::function
How to merge a list of samples to form a batch.pin_memory::logical(1)
Whether the dataloader copies tensors into CUDA pinned memory before returning them.drop_last::logical(1)
Whether to drop the last training batch in each epoch during training. Default isFALSE. It is ignored when abatch_sampleris provided.timeout::numeric(1)
The timeout value for collecting a batch from workers. Negative values mean no timeout and the default is-1.worker_init_fn::function(id)
A function that receives the worker id (in[1, num_workers]) and is executed after seeding on the worker but before data loading.worker_globals::list()|character()
When loading data in parallel, this allows to export globals to the workers. If this is a character vector, the objects in the global environment with those names are copied to the workers.worker_packages::character()
Which packages to load on the workers.
Also see torch::dataloader for more information.
Input and Output Channels
There is one input channel "input" that takes in ModelDescriptor during traing and a Task of the specified
task_type during prediction.
The output is NULL during training and a Prediction of given task_type during prediction.
State
A trained LearnerTorchModel.
Internals
A LearnerTorchModel is created by calling model_descriptor_to_learner() on the
provided ModelDescriptor that is received through the input channel.
Then the parameters are set according to the parameters specified in PipeOpTorchModel and
its $train() method is called on the Task stored in the ModelDescriptor.
Super classes
mlr3pipelines::PipeOp -> mlr3pipelines::PipeOpLearner -> PipeOpTorchModel
Methods
PipeOpTorchModel$new()
Creates a new instance of this R6 class.
Usage
PipeOpTorchModel$new(task_type, id = "torch_model", param_vals = list())Arguments
task_type(
character(1))
The task type of the model.id(
character(1))
Identifier of the resulting object.param_vals(
list())
List of hyperparameter settings, overwriting the hyperparameter settings that would otherwise be set during construction.