CrossEntropyLoss
-
class torch.nn.modules.loss.CrossEntropyLoss(weight=None, size_average=None, ignore_index=-100, reduce=None, reduction='mean', label_smoothing=0.0)[source] -
Этот критерий вычисляет функцию потерь кросс-энтропии между входными логитами и целевым значением.
Он полезен при обучении модели для решения задачи классификации с
Cклассами. Если указан необязательный аргументweight, он должен быть одномернымTensor, назначающим вес каждому из классов. Это особенно полезно при несбалансированном обучающем наборе.Ожидается, что
inputсодержит ненормализованные логиты для каждого класса (которые, как правило,notдолжны быть положительными или в сумме равняться 1).inputдолжен быть тензором размера для входных данных без батча, или при для случаяK-мерных данных. Последний вариант полезен для входных данных более высокой размерности, например при вычислении потерь кросс-энтропии для каждого пикселя 2D-изображений.Ожидается, что
target, используемое этим критерием, содержит одно из следующего:-
Индексы классов в диапазоне , где — количество классов; если задан
ignore_index, эта функция потерь также принимает этот индекс класса (индекс не обязательно должен принадлежать диапазону классов). Функцию потерь без редукции (то есть приreduction, равном'none') в этом случае можно описать следующим образом:где — входные данные, — целевое значение, — вес, — количество классов, а охватывает измерение мини-батча, а также для случая
K-мерных данных. Еслиreductionне равно'none'(значение по умолчанию —'mean'), тоОбратите внимание, что этот случай эквивалентен применению
LogSoftmaxк входным данным с последующим применениемNLLLoss. -
Вероятности для каждого класса; полезно, когда требуются метки, выходящие за рамки одного класса на элемент мини-батча, например при смешанных метках, сглаживании меток и т. д. Функцию потерь без редукции (то есть при
reduction, равном'none') в этом случае можно описать следующим образом:где — входные данные, — целевое значение, — вес, — количество классов, а охватывает измерение мини-батча, а также для случая
K-мерных данных. Еслиreductionне равно'none'(значение по умолчанию —'mean'), то
Примечание
Производительность этого критерия, как правило, выше, когда
targetсодержит индексы классов, поскольку это позволяет оптимизировать вычисления. Передавайтеtargetкак вероятности классов только в том случае, если одной метки класса на элемент мини-батча недостаточно.- Параметры:
-
-
weight (Tensor, optional) – вручную задаваемый вес масштабирования для каждого класса. Если задан, должен быть тензором размера
C. -
size_average (bool, optional) – Устарел (см.
reduction). По умолчанию потери усредняются по каждому элементу потерь в батче. Обратите внимание, что для некоторых функций потерь на один образец приходится несколько элементов. Если полеsize_averageустановлено вFalse, вместо этого потери суммируются для каждого мини-батча. Игнорируется, еслиreduceравноFalse. Значение по умолчанию:True -
ignore_index (int, optional) – задаёт целевое значение, которое игнорируется и не влияет на градиент входных данных. Если
size_averageравноTrue, потери усредняются по целевым значениям, которые не игнорируются. Обратите внимание, чтоignore_indexприменим только в том случае, если целевые значения содержат индексы классов. -
reduce (bool, optional) – Устарел (см.
reduction). По умолчанию потери усредняются или суммируются по наблюдениям каждого мини-батча в зависимости отsize_average. ЕслиreduceравноFalse, вместо этого возвращает значение потерь для каждого элемента батча и игнорируетsize_average. Значение по умолчанию:True -
reduction (str, optional) – задаёт редукцию, применяемую к выходным данным:
'none'|'mean'|'sum'.'none': редукция не применяется,'mean': вычисляется среднее значение выходных данных с учётом весов,'sum': выходные данные суммируются. Примечание:size_averageиreduceпостепенно выводятся из употребления; пока что указание любого из этих аргументов переопределитreduction. Значение по умолчанию:'mean' - label_smoothing (float, optional) – число с плавающей запятой в диапазоне [0.0, 1.0]. Задаёт степень сглаживания при вычислении потерь, где 0.0 означает отсутствие сглаживания. Целевые значения становятся смесью исходной истинной разметки и равномерного распределения, как описано в статье Переосмысление архитектуры Inception для компьютерного зрения. Значение по умолчанию: .
-
weight (Tensor, optional) – вручную задаваемый вес масштабирования для каждого класса. Если задан, должен быть тензором размера
- Форма:
-
- Входные данные: форма , или , где для функции потерь
K-мерной размерности. - Целевые значения: если содержат индексы классов, форма , или , где для функции потерь K-мерной размерности; каждое значение должно находиться в диапазоне . При использовании индексов классов тип данных целевых значений должен быть long. Если целевые значения содержат вероятности классов, их форма должна совпадать с формой входных данных, а каждое значение должно находиться в диапазоне . Это означает, что при использовании вероятностей классов тип данных целевых значений должен быть float. Обратите внимание, что PyTorch не проверяет строго ограничения на вероятности классов; пользователь должен убедиться, что
targetсодержит корректные распределения вероятностей (подробности приведены ниже в разделе с примерами). - Выходные данные: если reduction равно ‘none’, форма , или , где для функции потерь K-мерной размерности; форма зависит от формы входных данных. В остальных случаях — скаляр.
где:
- Входные данные: форма , или , где для функции потерь
Примеры
>>> # Example of target with class indices >>> loss = nn.CrossEntropyLoss() >>> input = torch.randn(3, 5, requires_grad=True) >>> target = torch.empty(3, dtype=torch.long).random_(5) >>> output = loss(input, target) >>> output.backward() >>> >>> # Example of target with class probabilities >>> input = torch.randn(3, 5, requires_grad=True) >>> target = torch.randn(3, 5).softmax(dim=1) >>> output = loss(input, target) >>> output.backward()
Примечание
Если
targetсодержит вероятности классов, он должен состоять из мягких меток — то есть каждая записьtargetдолжна представлять распределение вероятностей возможных классов для заданного образца данных, при этом отдельные вероятности должны находиться в диапазоне от[0,1], а сумма распределения должна равняться 1. Поэтому в приведённом выше примере вероятностей классов кtargetприменяется функцияsoftmax().PyTorch не проверяет, находятся ли значения, переданные в
target, в диапазоне[0,1], и равна ли сумма распределения для каждого образца данных1. Предупреждение не выводится; пользователь должен убедиться, чтоtargetсодержит корректные распределения вероятностей. Произвольные значения могут привести к вводящим в заблуждение значениям потерь и нестабильным градиентам во время обучения.Примеры
>>> # Example of target with incorrectly specified class probabilities >>> loss = nn.CrossEntropyLoss() >>> torch.manual_seed(283) >>> input = torch.randn(3, 5, requires_grad=True) >>> target = torch.randn(3, 5) >>> # Provided target class probabilities are not in range [0,1] >>> target tensor([[ 0.7105, 0.4446, 2.0297, 0.2671, -0.6075], [-1.0496, -0.2753, -0.3586, 0.9270, 1.0027], [ 0.7551, 0.1003, 1.3468, -0.3581, -0.9569]]) >>> # Provided target class probabilities do not sum to 1 >>> target.sum(axis=1) tensor([2.8444, 0.2462, 0.8873]) >>> # No error message and possible misleading loss value >>> loss(input, target).item() 4.6379876136779785 >>> >>> # Example of target with correctly specified class probabilities >>> # Use .softmax() to ensure true probability distribution >>> target_new = target.softmax(dim=1) >>> # New target class probabilities all in range [0,1] >>> target_new tensor([[0.1559, 0.1195, 0.5830, 0.1000, 0.0417], [0.0496, 0.1075, 0.0990, 0.3579, 0.3860], [0.2607, 0.1355, 0.4711, 0.0856, 0.0471]]) >>> # New target class probabilities sum to 1 >>> target_new.sum(axis=1) tensor([1.0000, 1.0000, 1.0000]) >>> loss(input, target_new).item() 2.55349063873291-
forward(input, target)[source] -
Выполняет прямой проход.
- Тип возвращаемого значения:
-
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.modules.loss.CrossEntropyLoss.html