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

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