Класс MultiLabelSoftMarginLoss
Класс MultiLabelSoftMarginLoss предназначен для задач многоклассовой классификации,
когда каждый объект может принадлежать одновременно нескольким классам.
Функция потерь принимает на вход два тензора: предсказания модели (не нормализованные логиты)
и целевые метки (бинарные векторы). Первым параметром конструктор принимает аргумент
reduction - способ агрегации потерь (по умолчанию 'mean'),
вторым - weight - веса для классов.
Синтаксис
loss = torch.nn.MultiLabelSoftMarginLoss(reduction='mean')
res = loss(pred, target)
Пример
Базовое использование функции потерь для двух образцов и трёх классов:
import torch
import torch.nn as nn
torch.manual_seed(0)
loss = nn.MultiLabelSoftMarginLoss()
pred = torch.tensor([
[0.2, -0.5, 1.3],
[1.2, 0.8, -0.1],
])
target = torch.tensor([
[1.0, 0.0, 1.0],
[1.0, 1.0, 0.0],
])
res = loss(pred, target)
print(res)
Результат выполнения кода:
tensor(0.6734)
Пример
Использование с параметром reduction='sum' для получения суммы потерь:
import torch
import torch.nn as nn
torch.manual_seed(0)
loss = nn.MultiLabelSoftMarginLoss(reduction='sum')
pred = torch.tensor([
[0.2, -0.5, 1.3],
[1.2, 0.8, -0.1],
])
target = torch.tensor([
[1.0, 0.0, 1.0],
[1.0, 1.0, 0.0],
])
res = loss(pred, target)
print(res)
Результат выполнения кода:
tensor(1.3468)
Пример
Использование весов классов для учета дисбаланса:
import torch
import torch.nn as nn
torch.manual_seed(0)
weights = torch.tensor([0.5, 2.0, 1.5])
loss = nn.MultiLabelSoftMarginLoss(weight=weights)
pred = torch.tensor([
[0.2, -0.5, 1.3],
[1.2, 0.8, -0.1],
])
target = torch.tensor([
[1.0, 0.0, 1.0],
[1.0, 1.0, 0.0],
])
res = loss(pred, target)
print(res)
Результат выполнения кода:
tensor(0.7705)
Пример
Сравнение с использованием BCEWithLogitsLoss для аналогичной задачи:
import torch
import torch.nn as nn
torch.manual_seed(0)
margin_loss = nn.MultiLabelSoftMarginLoss()
bce_loss = nn.BCEWithLogitsLoss()
pred = torch.tensor([
[0.2, -0.5, 1.3],
[1.2, 0.8, -0.1],
])
target = torch.tensor([
[1.0, 0.0, 1.0],
[1.0, 1.0, 0.0],
])
res_margin = margin_loss(pred, target)
res_bce = bce_loss(pred, target)
print(res_margin)
print(res_bce)
Результат выполнения кода:
tensor(0.6734)
tensor(0.4358)
Смотрите также
-
класс
BCEWithLogitsLoss,
который объединяет сигмоиду и бинарную кросс-энтропию -
класс
BCELoss,
который вычисляет бинарную кросс-энтропию -
класс
CrossEntropyLoss,
который подходит для задач классификации с одним классом -
класс
SoftMarginLoss,
который вычисляет потерю для бинарной классификации