MarginRankingLoss
-
class torch.nn.MarginRankingLoss(margin=0.0, size_average=None, reduce=None, reduction='mean')[source] -
Создаёт критерий, который измеряет потерю, заданную входными данными , , двумя 1D мини-пачками или 0D
Tensors, и меткой 1D мини-пачки или 0DTensor(содержащей 1 или -1).Если , то предполагается, что первый вход должен быть ранжирован выше (иметь большее значение), чем второй вход, и наоборот для .
Функция потерь для каждой пары выборок в мини-пачке:
- Параметры:
-
- margin (float, необязательно) – Значение по умолчанию .
-
size_average (bool, необязательно) – Устарело (см.
reduction). По умолчанию, потери усредняются по каждому элементу потерь в пакете. Обратите внимание, что для некоторых потерь может быть несколько элементов на образец. Если полеsize_averageустановлено в значениеFalse, потери вместо этого суммируются для каждой мини-пачки. Игнорируется, когдаreduceравноFalse. По умолчанию:True -
reduce (bool, необязательно) – Устарело (см.
reduction). По умолчанию, потери усредняются или суммируются по наблюдениям для каждой мини-пачки в зависимости отsize_average. КогдаreduceравноFalse, возвращает потерю на элемент пачки и игнорируетsize_average. По умолчанию:True -
reduction (str, необязательно) – Указывает на применение уменьшения к выходу:
'none'|'mean'|'sum'.'none': уменьшение не применяется,'mean': сумма вывода делится на количество элементов в выводе,'sum': вывод суммируется. Примечание:size_averageиreduceв процессе устаревания, и пока что указание любого из этих двух аргументов переопределитreduction. По умолчанию:'mean'
- Форма:
-
- Ввод1: или , где
N— размер пачки. - Ввод2: или , та же форма, что и у Ввода1.
- Цель: или , та же форма, что и у входов.
- Вывод: скаляр. Если
reduction—'none', и размер ввода не , то .
- Ввод1: или , где
Примеры:
>>> loss = nn.MarginRankingLoss() >>> input1 = torch.randn(3, requires_grad=True) >>> input2 = torch.randn(3, requires_grad=True) >>> target = torch.randn(3).sign() >>> output = loss(input1, input2, target) >>> output.backward()
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.nn.MarginRankingLoss.html