Класс CTCLoss
Класс CTCLoss реализует функцию потерь Connectionist Temporal
Classification (CTC), которая применяется в задачах распознавания
последовательностей, таких как распознавание рукописного текста,
распознавание речи или OCR. Основное преимущество CTC - возможность
обучения модели без необходимости точного выравнивания входных
данных с целевыми последовательностями.
При создании объекта класса можно указать несколько параметров:
blank - индекс пустого токена (по умолчанию 0),
reduction - способ агрегации потерь ('none', 'mean', 'sum'),
zero_infinity - флаг для замены бесконечных потерь на ноль.
Синтаксис
torch.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=False)
Параметры класса
Основные параметры конструктора CTCLoss:
-
blank(int) - индекс пустого токена, который будет игнорироваться при вычислении потерь. Значение по умолчанию - 0. -
reduction(str) - способ агрегации потерь: 'none' - возвращает потери для каждого элемента батча, 'mean' - усредняет потери по батчу, 'sum' - суммирует потери по батчу. -
zero_infinity(bool) - если установлено в True, то бесконечные потери заменяются на ноль. Полезно для стабильности обучения.
Входные параметры
Метод forward принимает три обязательных аргумента:
-
log_probs(Tensor) - логарифмы вероятностей, размерность (T, N, C) или (T, C) для одномерного батча, где T - длина последовательности, N - размер батча, C - количество классов (включая пустой токен). -
targets(Tensor) - целевые последовательности, размерность (N, S) или (S) для одномерного батча, где N - размер батча, S - максимальная длина целевой последовательности. -
input_lengths(Tensor) - длины входных последовательностей, размерность (N,). -
target_lengths(Tensor) - длины целевых последовательностей, размерность (N,).
Пример
Рассмотрим базовый пример использования CTCLoss для
распознавания последовательностей:
import torch
# Создаём функцию потерь
criterion = torch.nn.CTCLoss(blank=0, reduction='mean')
# Входные данные: логарифмы вероятностей
# T=3, N=2, C=4 (3 токена + пустой)
log_probs = torch.randn(3, 2, 4).log_softmax(2)
# Целевые последовательности
targets = torch.tensor([
[1, 2],
[1, 3],
])
# Длины последовательностей
input_lengths = torch.tensor([3, 3])
target_lengths = torch.tensor([2, 2])
# Вычисление потерь
loss = criterion(log_probs, targets, input_lengths, target_lengths)
print(loss)
Результат выполнения кода:
tensor(1.5823)
Пример
Использование CTCLoss с параметром zero_infinity для
предотвращения нестабильности при обучении:
import torch
# Создаём функцию потерь с заменой бесконечностей на ноль
criterion = torch.nn.CTCLoss(
blank=0,
reduction='sum',
zero_infinity=True
)
# Входные данные с экстремальными значениями
log_probs = torch.tensor([
[[-1000.0, -1000.0, -1000.0, -1000.0],
[-1000.0, -1000.0, -1000.0, -1000.0]],
[[-1000.0, -1000.0, -1000.0, -1000.0],
[-1000.0, -1000.0, -1000.0, -1000.0]],
[[-1000.0, -1000.0, -1000.0, -1000.0],
[-1000.0, -1000.0, -1000.0, -1000.0]]
])
targets = torch.tensor([1, 2])
input_lengths = torch.tensor([3, 3])
target_lengths = torch.tensor([1, 1])
loss = criterion(log_probs, targets, input_lengths, target_lengths)
print(loss)
Результат выполнения кода:
tensor(0.)
Пример
Пример использования CTCLoss с параметром reduction='none'
для получения потерь для каждого элемента батча:
import torch
# Создаём функцию потерь без агрегации
criterion = torch.nn.CTCLoss(reduction='none')
# Фиксируем случайность для воспроизводимости
torch.manual_seed(0)
# Генерируем данные
log_probs = torch.randn(4, 3, 5).log_softmax(2)
targets = torch.tensor([
[1, 2, 3],
[1, 3, 4],
[2, 3, 4],
])
input_lengths = torch.tensor([4, 4, 4])
target_lengths = torch.tensor([3, 3, 3])
# Вычисление потерь для каждого элемента
losses = criterion(log_probs, targets, input_lengths, target_lengths)
print(losses)
Результат выполнения кода:
tensor([2.1401, 1.6570, 1.7538])
Смотрите также
-
класс
CrossEntropyLoss,
который вычисляет кросс-энтропийную потерю для классификации -
класс
NLLLoss,
который вычисляет отрицательную логарифмическую правдоподобность -
класс
MSELoss,
который вычисляет среднеквадратичную ошибку -
класс
SmoothL1Loss,
который вычисляет сглаженную L1-потерю