Returns the function that converts the target column(s) of a task into the target tensor
y of a batch, i.e. the tensor that the loss is applied to.
The returned function takes an argument data, a data.table
containing only the target column(s), and returns a torch_tensor.
It is NULL for a task with no target at all, whose batches have no y element and whose loss
is called as loss(y_hat).
For the target encodings of the built-in task types, see section
Network Head and Target Encoding of LearnerTorch.
When adding support for a custom task type, implement a method for the corresponding
Task class.
Arguments
- task
(
Task)
The task.- ...
(any)
Additional arguments. Not used yet.
Examples
batchgetter = get_target_batchgetter(tsk("iris"))
batchgetter(data.table::data.table(Species = factor(c("setosa", "virginica"))))
#> torch_tensor
#> 1
#> 2
#> [ CPULongType{2} ]