Функция F.ctc_loss
Функция F.ctc_loss вычисляет потери CTC (Connectionist Temporal Classification), которые позволяют обучать модели распознаванию последовательностей без точного выравнивания входных и выходных данных. Это особенно полезно в задачах распознавания речи, обработки рукописного текста и других областях, где длина входной и выходной последовательностей может не совпадать. Функция принимает логиты модели, целевые метки, длины входных последовательностей и длины целевых последовательностей.
Синтаксис
torch.nn.functional.ctc_loss(log_probs, targets, input_lengths, target_lengths, blank=0, reduction='mean', zero_infinity=False)
Основные параметры функции:
-
log_probs- тензор размера (T, N, C) или (T, N, C) где T - длина временных шагов, N - размер батча, C - количество классов (включая blank); -
targets- тензор целевых меток размера (N, S) или (sum(target_lengths),), где S - максимальная длина целевой последовательности; -
input_lengths- тензор длин входных последовательностей для каждого элемента батча; -
target_lengths- тензор длин целевых последовательностей для каждого элемента батча; -
blank- индекс пустого символа (по умолчанию 0); -
reduction- способ редукции потерь ('none', 'mean', 'sum'); -
zero_infinity- заменять ли бесконечные потери на ноль.
Пример
Рассмотрим базовый пример вычисления CTC-потерь для батча из двух последовательностей:
import torch
import torch.nn.functional as F
torch.manual_seed(0)
# Параметры: временные шаги, батч, классы
T = 3
N = 2
C = 5
# Логиты (до применения softmax)
logits = torch.randn(T, N, C)
log_probs = F.log_softmax(logits, dim=2)
# Целевые метки: (N, S) - S - макс. длина цели
targets = torch.tensor([
[1, 2, 3],
[1, 2, 0]
])
# Длины входных последовательностей
input_lengths = torch.tensor([3, 3])
# Длины целевых последовательностей
target_lengths = torch.tensor([3, 2])
loss = F.ctc_loss(log_probs, targets, input_lengths, target_lengths)
print(loss)
Результат выполнения кода:
tensor(3.8795)
Пример
Пример с использованием параметра blank (индекс пустого символа) и других значений редукции:
import torch
import torch.nn.functional as F
torch.manual_seed(1)
T = 4
N = 1
C = 4
logits = torch.randn(T, N, C)
log_probs = F.log_softmax(logits, dim=2)
targets = torch.tensor([
[2, 3, 1, 0]
])
input_lengths = torch.tensor([4])
target_lengths = torch.tensor([3])
# Указываем blank=1 (по умолчанию 0)
loss_none = F.ctc_loss(log_probs, targets, input_lengths, target_lengths, blank=1, reduction='none')
loss_mean = F.ctc_loss(log_probs, targets, input_lengths, target_lengths, blank=1, reduction='mean')
loss_sum = F.ctc_loss(log_probs, targets, input_lengths, target_lengths, blank=1, reduction='sum')
print(loss_none)
print(loss_mean)
print(loss_sum)
Результат выполнения кода:
tensor([4.6950])
tensor(4.6950)
tensor(4.6950)
Пример
Пример использования zero_infinity для обработки случаев, когда потери могут быть бесконечными (например, при очень длинных входных последовательностях):
import torch
import torch.nn.functional as F
torch.manual_seed(2)
T = 2
N = 1
C = 3
logits = torch.randn(T, N, C)
log_probs = F.log_softmax(logits, dim=2)
# Целевая последовательность длиннее, чем входная
targets = torch.tensor([
[1, 2, 0]
])
input_lengths = torch.tensor([2])
target_lengths = torch.tensor([3])
# По умолчанию zero_infinity=False, получим inf
loss_inf = F.ctc_loss(log_probs, targets, input_lengths, target_lengths, zero_infinity=False)
# С zero_infinity=True получим 0
loss_zero = F.ctc_loss(log_probs, targets, input_lengths, target_lengths, zero_infinity=True)
print(loss_inf)
print(loss_zero)
Результат выполнения кода:
tensor(inf)
tensor(0.)
Смотрите также
-
функцию
cross_entropy,
которая вычисляет кросс-энтропийные потери -
функцию
nll_loss,
которая вычисляет отрицательное логарифмическое правдоподобие -
функцию
kl_div,
которая вычисляет дивергенцию Кульбака-Лейблера -
функцию
mse_loss,
которая вычисляет среднеквадратичную ошибку