Класс MultiMarginLoss
Класс MultiMarginLoss реализует многоклассовую маржинальную функцию
потерь, которая используется в задачах классификации с несколькими
классами. Эта функция потерь оптимизирует разделение между правильным
классом и остальными классами с заданным отступом. Первым параметром
конструктор принимает размер отступа, вторым параметром - размер весов,
третьим - способ редукции, четвёртым - режим суммирования.
Синтаксис
torch.nn.MultiMarginLoss(p=1, margin=1.0, weight=None, reduction='mean')
Параметры
Класс принимает следующие параметры:
-
p- степень нормы (обычно 1 или 2), по умолчанию 1 -
margin- размер отступа между правильным и неправильными классами, по умолчанию 1.0 -
weight- тензор весов для каждого класса, по умолчанию None -
reduction- способ редукции потерь ('none', 'mean', 'sum'), по умолчанию 'mean'
Пример использования с одномерными данными
Давайте создадим модель с одним линейным слоем и применим многоклассовую маржинальную потерю для трёх классов:
import torch
import torch.nn as nn
torch.manual_seed(0)
model = nn.Linear(10, 3)
criterion = nn.MultiMarginLoss()
x = torch.randn(1, 10)
y = torch.tensor([0])
output = model(x)
loss = criterion(output, y)
print(loss.item())
Результат выполнения кода:
0.8030233979225159
Пример с использованием весов классов
Рассмотрим пример с неравномерным распределением классов, где мы задаём веса для каждого класса:
import torch
import torch.nn as nn
torch.manual_seed(0)
weights = torch.tensor([0.5, 1.0, 2.0])
criterion = nn.MultiMarginLoss(weight=weights)
x = torch.randn(1, 10)
y = torch.tensor([1])
model = nn.Linear(10, 3)
output = model(x)
loss = criterion(output, y)
print(loss.item())
Результат выполнения кода:
0.4589586853981018
Пример с батчем данных
Применим функцию потерь к батчу из нескольких примеров и используем редукцию 'sum':
import torch
import torch.nn as nn
torch.manual_seed(0)
criterion = nn.MultiMarginLoss(reduction='sum')
x = torch.randn(4, 10)
y = torch.tensor([0, 2, 1, 0])
model = nn.Linear(10, 3)
output = model(x)
loss = criterion(output, y)
print(loss.item())
Результат выполнения кода:
25.771223068237305
Пример с отключением редукции
Используем редукцию 'none' для получения значений потерь для каждого примера в батче:
import torch
import torch.nn as nn
torch.manual_seed(0)
criterion = nn.MultiMarginLoss(reduction='none')
x = torch.randn(4, 10)
y = torch.tensor([0, 2, 1, 0])
model = nn.Linear(10, 3)
output = model(x)
loss = criterion(output, y)
print(loss)
Результат выполнения кода:
tensor([1.4286, 4.8443, 1.9691, 1.4108])
Смотрите также
-
класс
CrossEntropyLoss,
который вычисляет кросс-энтропийную потерю для задач классификации -
класс
MSELoss,
который вычисляет среднеквадратичную ошибку для задач регрессии -
класс
HingeEmbeddingLoss,
который вычисляет шарнирную потерю для задач эмбеддинга -
класс
TripletMarginLoss,
который вычисляет триплетную маржинальную потерю для обучения эмбеддингов