mapping_networks.losses.LossOutput

class mapping_networks.losses.LossOutput(total: Tensor, components: dict[str, ~torch.Tensor]=<factory>, metrics: dict[str, float]=<factory>)

Structured output from a loss computation.

total

Weighted scalar loss for backward().

Type:

torch.Tensor

components

Named component scalars (e.g. {"task": ..., "stability": ...}). Each value is a weighted, autograd-live scalar.

Type:

dict[str, torch.Tensor]

metrics

Unweighted scalar floats for logging and diagnostics. These are detached from the graph and safe to serialize.

Type:

dict[str, float]

__init__(total: Tensor, components: dict[str, ~torch.Tensor]=<factory>, metrics: dict[str, float]=<factory>) None

Methods

__init__(total, components, ...)

Attributes