WeightedMSELoss#

class core.nn.loss.WeightedMSELoss(*args, **kwargs)[source]#

Bases: torch.nn.modules.module.Module

Mean squared error under per-element weights, normalized by their sum.

forward(input, target, weights)[source]#

Compute mean squared error loss.

Parameters:
  • input (Tensor) – \((N, D)\) float, model predictions.

  • target (Tensor) – \((N, D)\) float, ground truth targets.

  • weights (Tensor) – \((N, D)\) float, per-element loss weights.

Return type:

Tensor