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 a mask of 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.