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

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