tf.nn.ctc_loss
Вычисляет потерю CTC (Connectionist Temporal Classification).
tf.nn.ctc_loss(
labels,
logits,
label_length,
logit_length,
logits_time_major=True,
unique=None,
blank_index=None,
name=None
)
Этот оператор реализует потерю CTC, как представлено в Graves et al., 2006.
Connectionist Temporal Classification (CTC) — это тип выходных данных нейронной сети и связанной функции оценки для обучения рекуррентных нейронных сетей (RNN), таких как LSTM-сети, для решения задач с изменяемым временным интервалом. Его можно использовать для задач, таких как распознавание рукописного ввода в режиме онлайн или распознавание фонем в аудиозаписи речи. CTC относится к выходным данным и оценке и не зависит от структуры базовой нейронной сети.
Примечания:
- Этот класс выполняет операцию softmax для вас, поэтому
logits, например, линейные проекции выходов LSTM. - Выходные данные показывают повторяющиеся классы с пропусками между ними и также могут выдавать повторяющиеся классы без пробелов, которые необходимо сгруппировать декодером.
-
labelsможет быть предоставлен либо в виде плотного, нулевого заполненияTensorс вектором длин последовательностей меток, либо какSparseTensor. - На TPU: Поддерживаются только плотные заполненные
labels. - На CPU и GPU: вызывающий код может использовать
SparseTensorили плотно заполненныеlabels, но вызов с использованиемSparseTensorбудет значительно быстрее. - По умолчанию метка пропуска —
0, а неnum_labels - 1(гдеnum_labels— размер внутреннего измеренияlogits), если не указано иное вblank_index.
tf.random.set_seed(50)
batch_size = 8
num_labels = 6
max_label_length = 5
num_frames = 12
labels = tf.random.uniform([batch_size, max_label_length],
minval=1, maxval=num_labels, dtype=tf.int64)
logits = tf.random.uniform([num_frames, batch_size, num_labels])
label_length = tf.random.uniform([batch_size], minval=2,
maxval=max_label_length, dtype=tf.int64)
label_mask = tf.sequence_mask(label_length, maxlen=max_label_length,
dtype=label_length.dtype)
labels *= label_mask
logit_length = [num_frames] * batch_size
with tf.GradientTape() as t:
t.watch(logits)
ref_loss = tf.nn.ctc_loss(
labels=labels,
logits=logits,
label_length=label_length,
logit_length=logit_length,
blank_index=0)
ref_grad = t.gradient(ref_loss, logits)| Аргументы | |
|---|---|
labels | Tensor формы [batch_size, max_label_seq_length] или SparseTensor. |
logits | Tensor формы [frames, batch_size, num_labels]. Если logits_time_major == False, форма — [batch_size, frames, num_labels]. |
label_length | Tensor формы [batch_size]. None, если labels является SparseTensor. Длина последовательности меток-ссылок в labels. |
logit_length | Tensor формы [batch_size]. Длина входной последовательности в logits. |
logits_time_major | (необязательно) Если True (по умолчанию), форма logits — [кадры, размер_пакета, количество_меток]. В противном случае форма — [batch_size, frames, num_labels]. |
unique | (необязательно) Уникальные индексы меток, вычисленные с помощью ctc_unique_labels(labels). Если предоставлены, включают более быструю и экономичную по памяти реализацию на TPU. |
blank_index | (необязательно) Установите индекс класса, используемый для метки пропуска. Отрицательные значения начнутся с num_labels, т.е. -1 воссоздаст поведение ctc_loss, используя num_labels - 1 для символа пропуска. Переключение с 0 по умолчанию влечет за собой некоторые затраты на память/производительность, так как может быть создана дополнительная сдвинутая копия logits. |
name | Имя этого Op. По умолчанию "ctc_loss_dense". |
| Возвращаемые значения | |
|---|---|
loss | 1-мерный float Tensor формы [batch_size], содержащий отрицательные логарифмические вероятности. |
| Исключения | |
|---|---|
ValueError | Аргумент blank_index должен быть предоставлен, когда labels — это SparseTensor. |
| Ссылки | |
|---|---|
| Connectionist Temporal Classification - Labeling Unsegmented Sequence Data with Recurrent Neural Networks: Graves et al., 2006 (pdf) https://en.wikipedia.org/wiki/Connectionist_temporal_classification |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/api_docs/python/tf/nn/ctc_loss