tf.compat.v1.nn.ctc_loss
Вычисляет потерю CTC (Connectionist Temporal Classification).
tf.compat.v1.nn.ctc_loss(
labels,
inputs=None,
sequence_length=None,
preprocess_collapse_repeated=False,
ctc_merge_repeated=True,
ignore_longer_outputs_than_inputs=False,
time_major=True,
logits=None
)
Этот оператор реализует потерю CTC, как представлено в (Graves et al., 2006).
Требования к входным данным:
sequence_length(b) <= time for all b max(labels.indices(labels.indices[:, 1] == b, 2)) <= sequence_length(b) for all b.
Примечания:
Этот класс выполняет операцию softmax за вас, поэтому входы должны быть, например, линейными проекциями выходов LSTM.
Размер внутреннего измерения тензора inputs, num_classes, представляет num_labels + 1 классов, где num_labels — это количество истинных меток, а наибольшее значение (num_classes - 1) зарезервировано для метки пробела.
Например, для словаря, содержащего 3 метки [a, b, c], num_classes = 4, а индексация меток — это {a: 0, b: 1, c: 2, blank: 3}.
Относительно аргументов preprocess_collapse_repeated и ctc_merge_repeated:
Если preprocess_collapse_repeated равно True, то перед вычислением потери выполняется предобработка, при которой повторяющиеся метки, переданные в потерю, объединяются в одну метку. Это полезно, если метки обучения получены, например, из принудительных выравниваний, и, следовательно, имеют ненужные повторения.
Если ctc_merge_repeated установлено в False, то глубоко внутри вычисления CTC повторяющиеся метки, отличные от пробела, не будут объединяться и интерпретируются как отдельные метки. Это упрощенная (нестандартная) версия CTC.
Вот таблица ожидаемого поведения (приблизительно) первого порядка:
-
preprocess_collapse_repeated=False,ctc_merge_repeated=TrueКлассическое поведение CTC: Выдает истинные повторяющиеся классы с пробелами между ними и также может выдавать повторяющиеся классы без пробелов между ними, которые должны быть объединены декодером.
-
preprocess_collapse_repeated=True,ctc_merge_repeated=FalseНикогда не учится выдавать повторяющиеся классы, поскольку они сворачиваются в метках входных данных перед обучением.
-
preprocess_collapse_repeated=False,ctc_merge_repeated=FalseВыдает повторяющиеся классы с пробелами между ними, но, как правило, не требует от декодера объединять/сжимать повторяющиеся классы.
-
preprocess_collapse_repeated=True,ctc_merge_repeated=TrueНе протестировано. Вероятно, не научится выдавать повторяющиеся классы.
Вариант ignore_longer_outputs_than_inputs позволяет указать поведение CTCLoss при работе с последовательностями, которые имеют выходные данные длиннее входных. Если true, CTCLoss просто вернет нулевой градиент для этих элементов, в противном случае возвращается ошибка InvalidArgument, останавливающая обучение.
| Аргументы | |
|---|---|
labels | 1-мерный целочисленный тензор. int32 SparseTensor. labels.indices[i, :] == [b, t] означает, что labels.values[i] хранит идентификатор для (партія b, время t). labels.values[i] должны принимать значения в [0, num_labels). Дополнительные сведения см. в core/ops/ctc_ops.cc. |
inputs | 3-мерный тензор. float Tensor. Если time_major == False, это будет тензор формы: Tensor формы: [batch_size, max_time, num_classes]. Если time_major == True (по умолчанию), это будет тензор формы: Tensor формы: [max_time, batch_size, num_classes]. Логиты. |
sequence_length | 1-мерный целочисленный вектор, размер [batch_size]. Длины последовательностей. |
preprocess_collapse_repeated | Булево значение. По умолчанию: False. Если True, повторяющиеся метки сворачиваются перед вычислением CTC. |
ctc_merge_repeated | Булево значение. По умолчанию: True. |
ignore_longer_outputs_than_inputs | Булево значение. По умолчанию: False. Если True, последовательности с выходными данными, более длинными, чем входные, будут игнорироваться. |
time_major | Формат формы тензоров inputs. Если True, эти тензоры должны быть формы [max_time, batch_size, num_classes]. Если False, эти тензоры должны быть формы [batch_size, max_time, num_classes]. Использование time_major = True (по умолчанию) немного эффективнее, поскольку оно избегает транспонирования в начале вычисления ctc_loss. Однако большинство данных TensorFlow имеют формат "партия-главное", поэтому эта функция также принимает входные данные в формате "партия-главное". |
logits | Псевдоним для входных данных. |
| Возвращаемое значение | |
|---|---|
1-мерный вещественный тензор, размер [batch], содержащий отрицательные логарифмические вероятности. |
| Исключения | |
|---|---|
TypeError | если labels не является SparseTensor. |
| Ссылки | |
|---|---|
| Connectionist Temporal Classification - Labeling Unsegmented Sequence Data with Recurrent Neural Networks: Graves et al., 2006 (pdf) |
© 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/compat/v1/nn/ctc_loss