Функция nn.ctc_greedy_decoder
Функция nn.ctc_greedy_decoder выполняет жадное (greedy) декодирование
выхода модели CTC (Connectionist Temporal Classification). Первым параметром
функция принимает тензор логитов формы [batch_size, max_time, num_classes].
Вторым параметром передается тензор sequence_length с фактическими длинами
последовательностей в батче. Третьим параметром указывается булево значение
merge_repeated, которое определяет, нужно ли объединять повторяющиеся
классы в итоговой последовательности.
Функция возвращает кортеж из двух элементов: разреженный тензор
SparseTensor с декодированными индексами и тензор с оценкой
логарифмической вероятности (log-probability) для каждой последовательности.
Жадное декодирование выбирает на каждом временном шаге наиболее вероятный класс,
а затем удаляет повторы и символы blank.
Синтаксис
tf.nn.ctc_greedy_decoder(inputs, sequence_length, merge_repeated=True)
Пример
Давайте создадим простой пример с батчем из одной последовательности
длиной 5 и 3 классами, где класс 0 является blank:
import tensorflow as tf
inputs = tf.constant([[
[0.1, 0.6, 0.3],
[0.1, 0.6, 0.3],
[0.7, 0.2, 0.1],
[0.1, 0.2, 0.7],
[0.1, 0.9, 0.0]
]], dtype=tf.float32)
sequence_length = tf.constant([5], dtype=tf.int32)
decoded, log_prob = tf.nn.ctc_greedy_decoder(
inputs, sequence_length
)
res = tf.sparse.to_dense(decoded[0])
print(res)
print(log_prob)
Результат выполнения кода:
tf.Tensor([[1 2]], shape=(1, 2), dtype=int64)
tf.Tensor([-3.912023], shape=(1,), dtype=float32)
Пример
Давайте рассмотрим, как параметр merge_repeated влияет на результат.
Установим его в False, чтобы сохранить повторяющиеся классы:
import tensorflow as tf
inputs = tf.constant([[
[0.1, 0.6, 0.3],
[0.1, 0.6, 0.3],
[0.7, 0.2, 0.1],
[0.1, 0.2, 0.7],
[0.1, 0.9, 0.0]
]], dtype=tf.float32)
sequence_length = tf.constant([5], dtype=tf.int32)
decoded, log_prob = tf.nn.ctc_greedy_decoder(
inputs, sequence_length, merge_repeated=False
)
res = tf.sparse.to_dense(decoded[0])
print(res)
print(log_prob)
Результат выполнения кода:
tf.Tensor([[1 1 2]], shape=(1, 3), dtype=int64)
tf.Tensor([-3.912023], shape=(1,), dtype=float32)
Пример
Давайте обработаем батч из двух последовательностей разной длины:
import tensorflow as tf
inputs = tf.constant([
[
[0.1, 0.6, 0.3],
[0.7, 0.2, 0.1],
[0.1, 0.2, 0.7],
[0.1, 0.9, 0.0],
[0.1, 0.9, 0.0]
],
[
[0.7, 0.2, 0.1],
[0.1, 0.6, 0.3],
[0.1, 0.2, 0.7],
[0.1, 0.9, 0.0],
[0.1, 0.9, 0.0]
]
]], dtype=tf.float32)
sequence_length = tf.constant([5, 4], dtype=tf.int32)
decoded, log_prob = tf.nn.ctc_greedy_decoder(
inputs, sequence_length
)
res = tf.sparse.to_dense(decoded[0])
print(res)
print(log_prob)
Результат выполнения кода:
tf.Tensor(
[[1 2]
[0 1 2]], shape=(2, 3), dtype=int64
)
tf.Tensor([-4.3122135 -3.912023 ], shape=(2,), dtype=float32)
Смотрите также
-
функцию
ctc_loss,
которая вычисляет CTC-потери для задачи распознавания последовательностей -
функцию
ctc_beam_search_decoder,
которая выполняет beam search декодирование выхода CTC-модели -
функцию
softmax,
которая преобразует логиты в вероятности классов -
функцию
sparse_softmax_cross_entropy_with_logits,
которая вычисляет кросс-энтропию для разреженных меток