Класс MSELoss
Класс MSELoss вычисляет среднеквадратичную ошибку между
предсказанными значениями модели и целевыми значениями.
Это одна из самых популярных функций потерь для задач регрессии.
При создании объекта можно задать параметры reduction,
который определяет способ агрегации ошибки, и size_average
(устаревший параметр).
Синтаксис
torch.nn.MSELoss(size_average=None, reduce=None, reduction='mean')
Параметры
Параметр reduction определяет способ агрегации:
-
'none'- не выполнять агрегацию, возвращать поэлементную ошибку -
'mean'- возвращать среднее значение ошибки (используется по умолчанию) -
'sum'- возвращать сумму ошибок
Пример
Создадим функцию потерь и вычислим ошибку для двух тензоров:
import torch
criterion = torch.nn.MSELoss()
predictions = torch.tensor([1.0, 2.0, 3.0])
targets = torch.tensor([1.5, 2.5, 3.5])
loss = criterion(predictions, targets)
print(loss)
Результат выполнения кода:
tensor(0.2500)
Пример
Используем параметр reduction='sum' для получения суммы ошибок:
import torch
criterion = torch.nn.MSELoss(reduction='sum')
predictions = torch.tensor([1.0, 2.0, 3.0])
targets = torch.tensor([1.5, 2.5, 3.5])
loss = criterion(predictions, targets)
print(loss)
Результат выполнения кода:
tensor(0.7500)
Пример
При reduction='none' возвращается поэлементная ошибка:
import torch
criterion = torch.nn.MSELoss(reduction='none')
predictions = torch.tensor([1.0, 2.0, 3.0])
targets = torch.tensor([1.5, 2.5, 3.5])
loss = criterion(predictions, targets)
print(loss)
Результат выполнения кода:
tensor([0.2500, 0.2500, 0.2500])
Пример
Используем MSELoss в процессе обучения линейной модели:
import torch
torch.manual_seed(0)
model = torch.nn.Linear(1, 1)
criterion = torch.nn.MSELoss()
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
x = torch.tensor([[1.0], [2.0], [3.0], [4.0]])
y = torch.tensor([[2.0], [4.0], [6.0], [8.0]])
for epoch in range(100):
optimizer.zero_grad()
predictions = model(x)
loss = criterion(predictions, y)
loss.backward()
optimizer.step()
with torch.no_grad():
print(model.weight.item())
print(model.bias.item())
Результат выполнения кода:
2.000828266143799
2.9700433051491894e-05
Смотрите также
-
класс
L1Loss,
который вычисляет среднюю абсолютную ошибку -
класс
CrossEntropyLoss,
который используется для задач классификации -
класс
SmoothL1Loss,
который сочетает свойства L1 и L2Loss -
класс
HuberLoss,
который устойчив к выбросам в данных