Skip to contents

Context for training a torch learner. This is the - mostly read-only - information callbacks have access to through the argument ctx. For more information on callbacks, see CallbackSet.

Public fields

learner

(Learner)
The torch learner.

task_train

(Task)
The training task.

task_valid

(Task or NULL)
The validation task.

loader_train

(torch::dataloader)
The data loader for training.

loader_valid

(torch::dataloader)
The data loader for validation.

measures_train

(list() of Measures)
Measures used for training.

measures_valid

(list() of Measures)
Measures used for validation.

network

(torch::nn_module)
The torch network.

optimizer

(torch::optimizer)
The optimizer.

loss_fn

(torch::nn_module)
The loss function.

total_epochs

(integer(1))
The total number of epochs the learner is trained for.

last_scores_train

(named list() or NULL)
The scores from the last evaluated epoch, calculated over all rows that were trained on in that epoch. Names are the ids of the training measures. If LearnerTorch sets eval_freq different from 1, this is NULL in all epochs that don't evaluate the model.

last_scores_valid

(list())
The scores from the last evaluated epoch, calculated over the complete validation task. Names are the ids of the validation measures. If LearnerTorch sets eval_freq different from 1, this is NULL in all epochs that don't evaluate the model.

last_loss

(numeric(1))
The loss from the last trainings batch.

y_hat

(torch_tensor)
The network's output for the current batch, or its first element if the network returns more than one tensor. Provided for the most common case where a network returns a single tensor. For the full output, see y_hats.

y_hats

(torch_tensor or list())
The complete output of the network for the current batch, i.e. what the loss is applied to. This is a list() if the network returns more than one tensor and identical to y_hat otherwise.

epoch

(integer(1))
The current epoch.

step

(integer(1))
The current iteration, i.e. the index of the batch within the current epoch. This is reset to 0 at the beginning of every epoch, so a globally unique step is (epoch - 1) * length(loader_train) + step.

prediction_encoder

(function())
The learner's prediction encoder.

batch

(named list() of torch_tensors)
The current batch.

terminate

(logical(1))
If this field is set to TRUE at the end of an epoch, training stops.

callbacks

(named list() of CallbackSets)
The callbacks that are active during training, named by their ids. This allows a callback to access the state of the other callbacks, which is for example what CallbackSetCheckpoint does to save them.

device

(torch::torch_device)
The device.

Methods


ContextTorch$new()

Creates a new instance of this R6 class.

Usage

ContextTorch$new(
  learner,
  task_train,
  task_valid = NULL,
  loader_train,
  loader_valid = NULL,
  measures_train = NULL,
  measures_valid = NULL,
  network,
  optimizer,
  loss_fn,
  total_epochs,
  prediction_encoder,
  eval_freq = 1L,
  device
)

Arguments

learner

(Learner)
The torch learner.

task_train

(Task)
The training task.

task_valid

(Task or NULL)
The validation task.

loader_train

(torch::dataloader)
The data loader for training.

loader_valid

(torch::dataloader or NULL)
The data loader for validation.

measures_train

(list() of Measures or NULL)
Measures used for training. Default is NULL.

measures_valid

(list() of Measures or NULL)
Measures used for validation.

network

(torch::nn_module)
The torch network.

optimizer

(torch::optimizer)
The optimizer.

loss_fn

(torch::nn_module)
The loss function.

total_epochs

(integer(1))
The total number of epochs the learner is trained for.

prediction_encoder

(function())
The learner's prediction encoder. See section Inheriting of LearnerTorch.

eval_freq

(integer(1))
The evaluation frequency.

device

(character(1))
The device.


ContextTorch$clone()

The objects of this class are cloneable with this method.

Usage

ContextTorch$clone(deep = FALSE)

Arguments

deep

Whether to make a deep clone.