cayleypy.train.Loss
- class cayleypy.train.Loss[source]
Base class for losses used to train distance-estimating models.
Subclasses define the per-element loss in
elementwise(). Calling the loss computes those values and reduces them to a scalar, which can be backpropagated.Losses are not modules (
torch.nn.Module): they have no learnable parameters, so keeping them out of the module tree keeps checkpoints free of entries that mean nothing at inference time.Multi-output (Q-) models predict a distance for every generator, i.e. their predictions have shape
[batch_size, n_generators], and it is common that only some of these outputs are labeled - a state sampled on a random walk has a known target for the generator that produced it, while targets for the other generators are unknown. Passing amaskof the labeled outputs makes both the loss and the gradient see only those, which is what makes training on such sparse labels possible.- __init__()
Methods
__init__()elementwise(predictions, targets)Computes loss for every element separately, without reduction.
- abstractmethod elementwise(predictions: Tensor, targets: Tensor) Tensor[source]
Computes loss for every element separately, without reduction.
- Parameters:
predictions – Predicted distances.
targets – Target distances, of the same shape as
predictions.
- Returns:
Tensor of the same shape as
predictions, with loss for every element.