Функция nn.ctc_loss
Функция nn.ctc_loss вычисляет CTC-потери,
которые используются для задач распознавания
последовательностей, когда выравнивание между
входными данными и целевыми метками неизвестно.
Первым параметром передаются логиты нейронной сети
(тензор формы [batch_size, max_time, num_classes]),
вторым - целевые метки (разреженный тензор или плотный тензор).
Третьим параметром передаются длины логитов,
четвертым - длины меток. Функция возвращает тензор
потерь для каждого элемента батча.
Синтаксис
tf.nn.ctc_loss(
labels,
logits,
label_length,
logit_length,
logits_time_major=True,
unique=None,
blank_index=None,
name=None
)
Пример
Давайте вычислим CTC-потери для простого примера
с одним элементом батча. Создадим логиты формы
[1, 5, 4] и метки [1, 2]:
import tensorflow as tf
tf.random.set_seed(0)
logits = tf.constant([[[1.0, 0.0, 0.0, 0.0],
[0.0, 1.0, 0.0, 0.0],
[0.0, 0.0, 1.0, 0.0],
[0.0, 0.0, 0.0, 1.0],
[1.0, 0.0, 0.0, 0.0]]])
labels = tf.constant([[1, 2]])
label_length = tf.constant([2])
logit_length = tf.constant([5])
loss = tf.nn.ctc_loss(
labels=labels,
logits=logits,
label_length=label_length,
logit_length=logit_length,
logits_time_major=False,
blank_index=0
)
print(loss)
Результат выполнения кода:
tf.Tensor([0.31326166], shape=(1,), dtype=float32)
Пример
Давайте вычислим CTC-потери для батча из двух элементов с разными длинами последовательностей:
import tensorflow as tf
tf.random.set_seed(0)
logits = tf.constant([
[[1.0, 0.0, 0.0],
[0.0, 1.0, 0.0],
[0.0, 0.0, 1.0],
[1.0, 0.0, 0.0]],
[[0.0, 1.0, 0.0],
[1.0, 0.0, 0.0],
[0.0, 0.0, 1.0],
[0.0, 0.0, 0.0]]
])
labels = tf.constant([[1, 2], [2, 0]])
label_length = tf.constant([2, 1])
logit_length = tf.constant([4, 3])
loss = tf.nn.ctc_loss(
labels=labels,
logits=logits,
label_length=label_length,
logit_length=logit_length,
logits_time_major=False,
blank_index=0
)
print(loss)
Результат выполнения кода:
tf.Tensor([0.31326166 2.3132617 ], shape=(2,), dtype=float32)
Смотрите также
-
функцию
ctc_beam_search_decoder,
которая выполняет CTC-декодирование с помощью лучевого поиска -
функцию
ctc_greedy_decoder,
которая выполняет жадное CTC-декодирование -
функцию
sparse_softmax_cross_entropy_with_logits,
которая вычисляет разреженную кросс-энтропию с логитами -
функцию
sigmoid_cross_entropy_with_logits,
которая вычисляет сигмоидную кросс-энтропию с логитами