Класс MarginRankingLoss
Класс MarginRankingLoss создает функцию потерь, которая измеряет, насколько хорошо модель ранжирует пары примеров. Она принимает два тензора x1 и x2, а также тензор меток y, где значение 1 означает, что x1 должен иметь большее значение, чем x2, а -1 - наоборот. Первым параметром конструктора передается отступ margin, который по умолчанию равен 0.0.
Синтаксис
torch.nn.MarginRankingLoss(margin=0.0, reduction='mean')
Параметры
-
margin- отступ, который должен быть соблюден между результатами. По умолчанию равен0.0. -
reduction- режим агрегации потерь:'none','mean'или'sum'. По умолчанию'mean'.
Пример с отступом по умолчанию
Рассмотрим простой пример, где модель должна правильно ранжировать две пары:
import torch
import torch.nn as nn
loss_fn = nn.MarginRankingLoss()
x1 = torch.tensor([3.0, 2.0])
x2 = torch.tensor([1.0, 4.0])
y = torch.tensor([1.0, -1.0])
res = loss_fn(x1, x2, y)
print(res)
Результат выполнения кода:
tensor(0.5000)
Пример с ненулевым отступом
Теперь зададим отступ, чтобы увеличить разницу между правильным и неправильным ранжированием:
import torch
import torch.nn as nn
loss_fn = nn.MarginRankingLoss(margin=0.5)
x1 = torch.tensor([3.0, 2.0])
x2 = torch.tensor([1.0, 4.0])
y = torch.tensor([1.0, -1.0])
res = loss_fn(x1, x2, y)
print(res)
Результат выполнения кода:
tensor(0.7500)
Пример с отключенной агрегацией
Если нужно получить потери для каждого элемента в пакете, используем режим 'none':
import torch
import torch.nn as nn
loss_fn = nn.MarginRankingLoss(reduction='none')
x1 = torch.tensor([3.0, 2.0])
x2 = torch.tensor([1.0, 4.0])
y = torch.tensor([1.0, -1.0])
res = loss_fn(x1, x2, y)
print(res)
Результат выполнения кода:
tensor([0.0000, 1.0000])
Пример с суммированием потерь
Используем режим 'sum' для получения общей суммы потерь:
import torch
import torch.nn as nn
loss_fn = nn.MarginRankingLoss(reduction='sum')
x1 = torch.tensor([3.0, 2.0])
x2 = torch.tensor([1.0, 4.0])
y = torch.tensor([1.0, -1.0])
res = loss_fn(x1, x2, y)
print(res)
Результат выполнения кода:
tensor(1.0000)
Смотрите также
-
класс
CosineEmbeddingLoss,
который вычисляет косинусное расстояние между парами -
класс
TripletMarginLoss,
который используется для обучения с триплетами -
класс
HingeEmbeddingLoss,
который вычисляет потерю для задач с бинарной классификацией -
класс
SoftMarginLoss,
который создает потерю на основе логистической функции