CTCLoss
-
class torch.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=False)[source]
-
Потеря Connectionist Temporal Classification.
Вычисляет потерю между непрерывной (неразбитой) временной серией и целевой последовательностью. CTCLoss суммирует вероятность возможных выравниваний входных данных с целевыми, генерируя значение потери, которое является дифференцируемым по отношению к каждому узлу ввода. Выравнивание входных данных с целевыми предполагается «многие-к-одному», что ограничивает длину целевой последовательности таким образом, что она должна быть длине входных данных.
- Параметры
-
- blank (int, необязательно) – метка blank. По умолчанию .
-
reduction (str, необязательно) – Указывает операцию агрегации для применения к выходным данным:
'none'|'mean'|'sum'.'none': без агрегации,'mean': значения потерь выходов будут разделены на длины целевых данных, а затем вычислено среднее по батчу,'sum': значения потерь выходов будут просуммированы. По умолчанию:'mean' -
zero_infinity (bool, необязательно) – Нулевое значение бесконечных потерь и соответствующих градиентов. По умолчанию:
FalseБесконечные потери в основном возникают, когда входные данные слишком короткие, чтобы быть выровненными с целевыми данными.
- Форма:
-
- Log_probs: Тензор размера или , где , , и . Логарифмированные вероятности выходов (например, полученные с помощью
torch.nn.functional.log_softmax()). - Targets: Тензор размера или , где и . Представляет целевые последовательности. Каждый элемент в целевой последовательности — индекс класса. И индекс целевого значения не может быть blank (по умолчанию=0). В форме целевые данные дополняются до длины самой длинной последовательности и объединяются. В форме целевые данные предполагаются без заполнения и объединяются по одному измерению.
- Input_lengths: Кортеж или тензор размера или , где . Представляет длины входов (должны быть ). И длины указаны для каждой последовательности, чтобы обеспечить маскирование при предположении, что последовательности дополняются до одинаковых длин.
- Target_lengths: Кортеж или тензор размера или , где . Представляет длины целевых данных. Длины указаны для каждой последовательности, чтобы обеспечить маскирование при предположении, что последовательности дополняются до одинаковых длин. Если форма целевых данных , target_lengths фактически представляют собой индекс остановки для каждой целевой последовательности, такой что
target_n = targets[n,0:s_n]для каждой целевой последовательности в батче. Длины должны быть Если целевые данные заданы как 1-мерный тензор, являющийся конкатенацией отдельных целевых данных, то target_lengths должны суммироваться до общей длины тензора. - Выход: скаляр, если
reductionравно'mean'(по умолчанию) или'sum'. Еслиreductionравно'none', то , если вход является батчем, или , если вход не является батчем, где .
- Log_probs: Тензор размера или , где , , и . Логарифмированные вероятности выходов (например, полученные с помощью
Примеры:
>>> # Target are to be padded >>> T = 50 # Input sequence length >>> C = 20 # Number of classes (including blank) >>> N = 16 # Batch size >>> S = 30 # Target sequence length of longest target in batch (padding length) >>> S_min = 10 # Minimum target length, for demonstration purposes >>> >>> # Initialize random batch of input vectors, for *size = (T,N,C) >>> input = torch.randn(T, N, C).log_softmax(2).detach().requires_grad_() >>> >>> # Initialize random batch of targets (0 = blank, 1:C = classes) >>> target = torch.randint(low=1, high=C, size=(N, S), dtype=torch.long) >>> >>> input_lengths = torch.full(size=(N,), fill_value=T, dtype=torch.long) >>> target_lengths = torch.randint(low=S_min, high=S, size=(N,), dtype=torch.long) >>> ctc_loss = nn.CTCLoss() >>> loss = ctc_loss(input, target, input_lengths, target_lengths) >>> loss.backward() >>> >>> >>> # Target are to be un-padded >>> T = 50 # Input sequence length >>> C = 20 # Number of classes (including blank) >>> N = 16 # Batch size >>> >>> # Initialize random batch of input vectors, for *size = (T,N,C) >>> input = torch.randn(T, N, C).log_softmax(2).detach().requires_grad_() >>> input_lengths = torch.full(size=(N,), fill_value=T, dtype=torch.long) >>> >>> # Initialize random batch of targets (0 = blank, 1:C = classes) >>> target_lengths = torch.randint(low=1, high=T, size=(N,), dtype=torch.long) >>> target = torch.randint(low=1, high=C, size=(sum(target_lengths),), dtype=torch.long) >>> ctc_loss = nn.CTCLoss() >>> loss = ctc_loss(input, target, input_lengths, target_lengths) >>> loss.backward() >>> >>> >>> # Target are to be un-padded and unbatched (effectively N=1) >>> T = 50 # Input sequence length >>> C = 20 # Number of classes (including blank) >>> >>> # Initialize random batch of input vectors, for *size = (T,C) >>> input = torch.randn(T, C).log_softmax(1).detach().requires_grad_() >>> input_lengths = torch.tensor(T, dtype=torch.long) >>> >>> # Initialize random batch of targets (0 = blank, 1:C = classes) >>> target_lengths = torch.randint(low=1, high=T, size=(), dtype=torch.long) >>> target = torch.randint(low=1, high=C, size=(target_lengths,), dtype=torch.long) >>> ctc_loss = nn.CTCLoss() >>> loss = ctc_loss(input, target, input_lengths, target_lengths) >>> loss.backward()
- Ссылка:
-
A. Graves et al.: Connectionist Temporal Classification: Labelling Unsegmented Sequence Data with Recurrent Neural Networks: https://www.cs.toronto.edu/~graves/icml_2006.pdf
Примечание
Для использования CuDNN необходимо выполнить следующие условия:
targetsдолжно быть в конкатенированном формате, всеinput_lengthsдолжны бытьT. ,target_lengths, целочисленные аргументы должны иметь типtorch.int32.Регулярная реализация использует тип данных (более распространённый в PyTorch)
torch.long.Примечание
В некоторых случаях при использовании CUDA-бекенда с CuDNN данный оператор может выбрать недетерминированный алгоритм для повышения производительности. Если это нежелательно, вы можете попытаться сделать операцию детерминированной (возможно, с потерей производительности) путём установки
torch.backends.cudnn.deterministic = True. Для справки ознакомьтесь с примечаниями о Воспроизводимости.
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.nn.CTCLoss.html