torch.nn.functional.ctc_loss
-
torch.nn.functional.ctc_loss(log_probs, targets, input_lengths, target_lengths, blank=0, reduction='mean', zero_infinity=False)[source] -
Потеря классификации временных последовательностей с использованием свёрточных нейронных сетей.
Подробности см. в
CTCLoss.Примечание
В некоторых случаях, при использовании тензоров на устройстве CUDA и CuDNN, этот оператор может выбрать недетерминированный алгоритм для повышения производительности. Если это нежелательно, вы можете попытаться сделать операцию детерминированной (возможно, за счёт производительности) с помощью
torch.backends.cudnn.deterministic = True. Дополнительная информация по этому вопросу представлена в Воспроизводимости.Примечание
Этот оператор может генерировать недетерминированные градиенты при использовании тензоров на устройстве CUDA. Дополнительная информация по этому вопросу представлена в Воспроизводимости.
- Параметры:
-
-
log_probs (Тензор) – или , где
C = number of characters in alphabet including blank,T = input length, иN = batch size. Вероятности выходных данных в логарифмическом виде (например, полученные с помощьюtorch.nn.functional.log_softmax()). -
targets (Тензор) – или
(sum(target_lengths)). Целевые значения не могут быть пустыми. Во втором формате целевые значения предполагаются конкатенированными. - input_lengths (Тензор) – или . Длины входных данных (должны быть каждый )
- target_lengths (Тензор) – или . Длины целевых значений
- blank (int, необязательно) – Значение пустого символа. По умолчанию .
-
reduction (str, необязательно) – Указывает способ уменьшения выходного значения:
'none'|'mean'|'sum'.'none': уменьшение не будет применено,'mean': значения потерь будут разделены на длину целевых значений, а затем вычисляется среднее по пакету,'sum': значения будут суммированы. По умолчанию:'mean' -
zero_infinity (bool, необязательно) – Признак нулевых бесконечных потерь и связанных с ними градиентов. По умолчанию:
FalseБесконечные потери чаще всего возникают, когда входные данные слишком короткие для выравнивания с целевыми значениями.
-
log_probs (Тензор) – или , где
- Тип возвращаемого значения:
Пример:
>>> log_probs = torch.randn(50, 16, 20).log_softmax(2).detach().requires_grad_() >>> targets = torch.randint(1, 20, (16, 30), dtype=torch.long) >>> input_lengths = torch.full((16,), 50, dtype=torch.long) >>> target_lengths = torch.randint(10,30,(16,), dtype=torch.long) >>> loss = F.ctc_loss(log_probs, targets, input_lengths, target_lengths) >>> loss.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.functional.ctc_loss.html