Класс CTC
Класс CTC вычисляет функцию потерь
Connectionist Temporal Classification (CTC).
Он применяется в задачах распознавания речи
и рукописного ввода, где входная и целевая
последовательности имеют разную длину и
не выровнены друг относительно друга.
Первым параметром конструктор принимает
словарь или строку с настройками, например
reduction для способа сокращения потерь.
Метод call принимает истинные метки
y_true, предсказанные логиты y_pred,
длины меток label_length и длины
предсказаний logit_length.
Синтаксис
tf.keras.losses.CTC(
reduction="sum_over_batch_size",
name="ctc"
)
Пример
Давайте вычислим CTC-потерю для одной
последовательности. Истинные метки зададим
тензором [1, 2], а логиты - тензором
размера 4 на 5:
import tensorflow as tf
loss = tf.keras.losses.CTC()
y_true = tf.constant([[1, 2, 0, 0]])
y_pred = tf.random.uniform((1, 4, 5))
label_length = tf.constant([2])
logit_length = tf.constant([4])
res = loss(y_true, y_pred, label_length, logit_length)
print(res)
Результат выполнения кода:
tf.Tensor(4.6399364, shape=(), dtype=float32)
Пример
Давайте вычислим CTC-потерю для батча из двух последовательностей:
import tensorflow as tf
tf.random.set_seed(0)
loss = tf.keras.losses.CTC()
y_true = tf.constant([
[1, 2, 0],
[3, 0, 0]
])
y_pred = tf.random.uniform((2, 5, 4))
label_length = tf.constant([2, 1])
logit_length = tf.constant([5, 5])
res = loss(y_true, y_pred, label_length, logit_length)
print(res)
Результат выполнения кода:
tf.Tensor(3.7216048, shape=(), dtype=float32)
Пример
Давайте получим потерю для каждого элемента
батча отдельно, указав reduction="none":
import tensorflow as tf
tf.random.set_seed(0)
loss = tf.keras.losses.CTC(reduction="none")
y_true = tf.constant([
[1, 2, 0],
[3, 0, 0]
])
y_pred = tf.random.uniform((2, 5, 4))
label_length = tf.constant([2, 1])
logit_length = tf.constant([5, 5])
res = loss(y_true, y_pred, label_length, logit_length)
print(res)
Результат выполнения кода:
tf.Tensor([3.5961368 3.8470726], shape=(2,), dtype=float32)
Смотрите также
-
класс
CategoricalCrossentropy,
который вычисляет категориальную кросс-энтропию -
класс
SparseCategoricalCrossentropy,
который вычисляет разреженную категориальную кросс-энтропию -
класс
BinaryCrossentropy,
который вычисляет бинарную кросс-энтропию -
класс
KLDivergence,
который вычисляет дивергенцию Кульбака-Лейблера