MSERegressionLoss#

class MSERegressionLoss(reduction='mean', loss_weight=1.0)[source]#

Bases: Module

Mean-squared-error (L2) regression loss for continuous targets.

forward(pred, target) returns loss_weight * mse_loss(pred, target); pred and target must broadcast. Registered as MSERegressionLoss – use in a criteria=[...] list or as a per-head criterion.

Parameters:
  • reduction (str) – Reduction mode ("mean", "sum", "none"). Defaults to "mean".

  • loss_weight (float) – Global scale on the returned loss. Defaults to 1.0.

Example

>>> import torch
>>> from pimm.models.losses.builder import build_criteria
>>> crit = build_criteria([dict(type="MSERegressionLoss", loss_weight=1.0)])
>>> crit(torch.zeros(4), torch.full((4,), 2.0))  # mean (0 - 2)^2 = 4.0
tensor(4.)
MSERegressionLoss.forward(pred, target)[source]#