meansquarederror
Metric
MeanSquaredErrorfromtorchmetrics(torchmetrics.MeanSquaredError)
When to invoke this skill
The user has predictions + ground truth and asks to evaluate with MeanSquaredError, or
mentions torchmetrics.MeanSquaredError directly, or wants the standard torchmetrics implementation.
Reference signature
from torchmetrics import MeanSquaredError
# MeanSquaredError(squared: bool = True, num_outputs: int = 1, **kwargs: Any) -> None
Library docstring
Compute `mean squared error`_ (MSE).
.. math:: \text{MSE} = \frac{1}{N}\sum_i^N(y_i - \hat{y_i})^2
Where :math:`y` is a tensor of target values, and :math:`\hat{y}` is a tensor of predictions.
As input to ``forward`` and ``update`` the metric accepts the following input:
- ``preds`` (:class:`~torch.Tensor`): Predictions from model
- ``target`` (:class:`~torch.Tensor`): Ground truth values
As output of ``forward`` and ``compute`` the metric returns the following output:
- ``mean_squared_error`` (:class:`~torch.Tensor`): A tensor with the mean squared error
Args:
squared: If True returns MSE value, if False returns RMSE value.
num_outputs: Number of outputs in multioutput setting
kwargs: Additional keyword arguments, see :ref:`Metric kwargs` for more info.
Example::
Single output mse computation:
>>> from torch import tensor
>>> from torchmetrics.regression import MeanSquaredError
>>> target = tensor([2.5, 5.0, 4.0, 8.0])
>>> preds = tensor([3.0, 5.0, 2.5, 7.0])
>>> mean_squared_error = MeanSquaredError()
>>> mean_squared_error(preds, target)
tensor(0.8750)
Example::
Multioutput mse computation:
>>> from torch import tensor
>>> from torchmetrics.regression import MeanSquaredError
>>> target = tensor([[0.0, 0.0, 0.0], [0.0, 0.0, 0.0]])
>>> preds = tensor([[1.0, 2.0, 3.0], [1.0, 2.0, 3.0]])
>>> mean_squared_error = MeanSquaredError(num_outputs=3)
>>> mean_squared_error(preds, target)
tensor([1., 4., 9.])
Quick recipe
import torchmetrics as _m
score = _m.MeanSquaredError(y_true, y_pred)
Don'ts
- Don't reimplement when the library version handles edge cases (NaN, ties, empty inputs) better than a hand-rolled formula.
- Always check the library version's argument order — sklearn is
(y_true, y_pred)while torchmetrics is(preds, target).