Функция nn.ctc_beam_search_decoder
Функция nn.ctc_beam_search_decoder выполняет
CTC (Connectionist Temporal Classification) декодирование
методом лучевого поиска. Первым параметром функция
принимает тензор логитов формы [batch_size, max_time, num_classes].
Вторым параметром передается целое число beam_width -
ширина луча поиска. Третьим необязательным параметром
можно передать top_paths - количество лучших путей
для возврата. Функция возвращает кортеж из двух элементов:
список разреженных тензоров с декодированными
последовательностями и список тензоров с логарифмами
вероятностей для каждого пути.
Синтаксис
tf.nn.ctc_beam_search_decoder(inputs, sequence_length, beam_width=100, top_paths=1, merge_repeated=True)
Пример
Давайте выполним CTC декодирование для тензора логитов с тремя временными шагами и четырьмя классами:
import tensorflow as tf
inputs = tf.constant([
[[0.1, 0.6, 0.1, 0.2],
[0.1, 0.1, 0.7, 0.1],
[0.1, 0.1, 0.1, 0.7]]
])
sequence_length = tf.constant([3])
decoded, log_probs = tf.nn.ctc_beam_search_decoder(
inputs, sequence_length, beam_width=10, top_paths=1
)
print(decoded)
print(log_probs)
Результат выполнения кода:
[<tf.SparseTensor: shape=(1, 3), dtype=int64>]
tf.Tensor([[-1.6613297]], shape=(1, 1), dtype=float32)
Пример
Давайте извлечем декодированную последовательность из разреженного тензора:
import tensorflow as tf
inputs = tf.constant([
[[0.1, 0.6, 0.1, 0.2],
[0.1, 0.1, 0.7, 0.1],
[0.1, 0.1, 0.1, 0.7]]
])
sequence_length = tf.constant([3])
decoded, log_probs = tf.nn.ctc_beam_search_decoder(
inputs, sequence_length, beam_width=10, top_paths=1
)
res = tf.sparse.to_dense(decoded[0])
print(res)
Результат выполнения кода:
tf.Tensor([[1 2 3]], shape=(1, 3), dtype=int64)
Пример
Давайте вернем несколько лучших путей,
указав top_paths равным 3:
import tensorflow as tf
inputs = tf.constant([
[[0.1, 0.6, 0.1, 0.2],
[0.1, 0.1, 0.7, 0.1],
[0.1, 0.1, 0.1, 0.7]]
])
sequence_length = tf.constant([3])
decoded, log_probs = tf.nn.ctc_beam_search_decoder(
inputs, sequence_length, beam_width=10, top_paths=3
)
for i in range(3):
res = tf.sparse.to_dense(decoded[i])
print(res.numpy(), log_probs[i].numpy())
Результат выполнения кода:
[[1 2 3]] [-1.6613297]
[[1 3]] [-2.0613296]
[[2 3]] [-2.3613296]
Смотрите также
-
функцию
ctc_loss,
которая вычисляет CTC потери -
функцию
ctc_greedy_decoder,
которая выполняет жадное CTC декодирование -
функцию
softmax,
которая применяет softmax к логитам -
функцию
log_softmax,
которая вычисляет логарифм softmax