РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
542 of 769 menu

Класс 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-потерю
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить