Converts the raw output of a network into a list() that can be passed to
mlr3::as_prediction_data(), which is what the private .encode_prediction() method of a
LearnerTorch has to return.
This is the default implementation that is used by LearnerTorch and
LearnerTorchModel, i.e. by all learners that don't overwrite
.encode_prediction().
When adding support for a custom task type, implement a method for the corresponding
Task class, which makes the generic torch learners work for that task type.
For the network output that is expected for the built-in task types, see section
Network Head and Target Encoding of LearnerTorch.
Arguments
- task
(
Task)
The task to predict on.- network_output
(
torch_tensororlist()of them)
The raw output of the network in evaluation mode. A network with more than one head – e.g. one predicting a mean and a standard deviation – returns alist()of tensors, which is passed on unchanged. The encodings of the built-in task types expect a single tensor.- predict_type
(
character(1))
The predict type of the learner, e.g."response"or"prob".- ...
(any)
Additional arguments. Not used yet.
Value
named list()