NLLLoss
-
class torch.nn.NLLLoss(weight=None, size_average=None, ignore_index=- 100, reduce=None, reduction='mean')[source] -
Функция потерь отрицательного логарифмического правдоподобия. Она полезна для обучения задач классификации с
Cклассами.Если задан, необязательный аргумент
weightдолжен быть одномерным тензором, присваивающим вес каждому из классов. Это особенно полезно, когда у вас есть несбалансированный обучающий набор.Значение
input, полученное через вызов forward, должно содержать логарифмы вероятностей каждого класса.inputдолжен быть тензором размера либо или с для многомерного случаяK. Последний вариант полезен для ввода большей размерности, например, для вычисления потери NLL на каждый пиксель для 2D изображений.Получение логарифмов вероятностей в нейронной сети легко достигается добавлением слоя
LogSoftmaxв последний слой вашей сети. Вы можете использоватьCrossEntropyLossвместо этого, если не хотите добавлять дополнительный слой.Значение
targetкоторое ожидает эта функция потерь, должно быть индексом класса в диапазоне гдеC = number of classes; еслиignore_indexуказан, эта функция потерь также принимает этот индекс класса (этот индекс может не обязательно находиться в диапазоне классов).Несведённая (т.е. с
reductionустановленным в'none') функция потерь может быть описана как:где — вход, — цель, — вес, а — размер пакета. Если
reductionне'none'(по умолчанию'mean'), то- Parameters:
-
-
weight (Tensor, optional) – ручной масштабирующий вес, присваиваемый каждому классу. Если задан, он должен быть тензором размера
C. В противном случае он обрабатывается так, как будто имеет все единицы. -
size_average (bool, optional) – устарело (см.
reduction). По умолчанию потери усредняются по каждому элементу потерь в пачке. Обратите внимание, что для некоторых потерь есть несколько элементов на образец. Если полеsize_averageустановлено вFalse, потери вместо этого суммируются для каждой мини-пачки. Игнорируется, когдаreduceравноFalse. По умолчанию:None -
ignore_index (int, optional) – указывает целевое значение, которое игнорируется и не вносит вклад в градиент входных данных. Когда
size_averageравноTrue, потери усредняются по неигнорируемым целям. -
reduce (bool, optional) – устарело (см.
reduction). По умолчанию потери усредняются или суммируются по наблюдениям для каждой мини-пачки в зависимости отsize_average. КогдаreduceравноFalse, возвращает потери на элемент каждой пачки и игнорируетsize_average. По умолчанию:None -
reduction (str, optional) – указывает снижение, которое нужно применить к выводу:
'none'|'mean'|'sum'.'none': не применяется снижение,'mean': берется взвешенное среднее значение выхода,'sum': вывод будет суммирован. Примечание:size_averageиreduceнаходятся в процессе устаревания, и пока что указание одного из этих двух аргументов переопределитreduction. По умолчанию:'mean'
-
weight (Tensor, optional) – ручной масштабирующий вес, присваиваемый каждому классу. Если задан, он должен быть тензором размера
- Форма:
-
- Вход: или , где
C = number of classes, или с в случаеK-мерной потери. - Цель: или , где каждое значение равно , или с в случае K-мерной потери.
- Выход: Если
reductionравно'none', форма или с в случае K-мерной потери. В противном случае - скаляр.
- Вход: или , где
Примеры:
>>> m = nn.LogSoftmax(dim=1) >>> loss = nn.NLLLoss() >>> # input is of size N x C = 3 x 5 >>> input = torch.randn(3, 5, requires_grad=True) >>> # each element in target has to have 0 <= value < C >>> target = torch.tensor([1, 0, 4]) >>> output = loss(m(input), target) >>> output.backward() >>> >>> >>> # 2D loss example (used, for example, with image inputs) >>> N, C = 5, 4 >>> loss = nn.NLLLoss() >>> # input is of size N x C x height x width >>> data = torch.randn(N, 16, 10, 10) >>> conv = nn.Conv2d(16, C, (3, 3)) >>> m = nn.LogSoftmax(dim=1) >>> # each element in target has to have 0 <= value < C >>> target = torch.empty(N, 8, 8, dtype=torch.long).random_(0, C) >>> output = loss(m(conv(data)), 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.NLLLoss.html