From 19ff7fba3d30e3658a21ca46fc8a1cd7d2342030 Mon Sep 17 00:00:00 2001 From: Roman Kazantsev Date: Mon, 21 Aug 2023 14:26:36 +0400 Subject: [PATCH] [TF FE] Fix support of CTCLoss and add tests to pre-commit (#19291) Signed-off-by: Kazantsev, Roman --- .../tensorflow_common/src/op/ctc_loss.cpp | 73 +++++++++++-------- .../tensorflow_tests/test_tf_CTCLoss.py | 30 ++++---- 2 files changed, 55 insertions(+), 48 deletions(-) diff --git a/src/frontends/tensorflow_common/src/op/ctc_loss.cpp b/src/frontends/tensorflow_common/src/op/ctc_loss.cpp index b4b72bb407f..1abba8801f2 100644 --- a/src/frontends/tensorflow_common/src/op/ctc_loss.cpp +++ b/src/frontends/tensorflow_common/src/op/ctc_loss.cpp @@ -2,14 +2,24 @@ // SPDX-License-Identifier: Apache-2.0 // +#include "openvino/op/ctc_loss.hpp" + #include "common_op_table.hpp" -#include "openvino/opsets/opset8.hpp" +#include "openvino/op/broadcast.hpp" +#include "openvino/op/constant.hpp" +#include "openvino/op/convert.hpp" +#include "openvino/op/equal.hpp" +#include "openvino/op/reduce_sum.hpp" +#include "openvino/op/scatter_nd_update.hpp" +#include "openvino/op/select.hpp" +#include "openvino/op/shape_of.hpp" +#include "openvino/op/slice.hpp" using namespace std; using namespace ov; -using namespace opset8; using namespace ov::frontend; -using namespace frontend::tensorflow; +using namespace ov::frontend::tensorflow; +using namespace ov::op; namespace ov { namespace frontend { @@ -17,7 +27,7 @@ namespace tensorflow { namespace op { OutputVector translate_ctc_loss_op(const NodeContext& node) { - // This is a translator for CTCLoss v1 aka tf.compat.v1.nn.ctc_loss + // this is a translator for CTCLoss v1 aka tf.compat.v1.nn.ctc_loss default_op_checks(node, 4, {"CTCLoss"}); auto logits = node.get_input(0); auto decoded_indices = node.get_input(1); @@ -34,40 +44,41 @@ OutputVector translate_ctc_loss_op(const NodeContext& node) { // we need to transpose it into [batch_size, time_size, num_classes] format // from [time_size, batch_size, num_classes] AxisVector logits_order = {1, 0, 2}; - logits = tensorflow::make_transpose(logits, logits_order); + logits = make_transpose(logits, logits_order); } - // Transform decoded labels from the sparse format into dense format - // Convert to the signed type since the mask with minus one is formed below - decoded_values = make_shared(decoded_values, element::i64); + // transform decoded labels from the sparse format into dense format + // convert to the signed type since the mask with minus one is formed below + decoded_values = make_shared(decoded_values, element::i64); // OpenVINO ScatterND operation requires indices to be signed - decoded_indices = make_shared(decoded_indices, element::i64); + decoded_indices = make_shared(decoded_indices, element::i64); // OpenVINO CTCLoss requires logit_length to be signed - logit_length = make_shared(logit_length, element::i64); + logit_length = make_shared(logit_length, element::i64); - auto logits_shape = make_shared(logits, element::i64); - auto dense_shape = make_shared(logits_shape, - make_shared(element::i64, Shape{}, 0), - make_shared(element::i64, Shape{}, 2), - make_shared(element::i64, Shape{}, 1)); - auto minus_one_value = make_shared(element::i64, Shape{}, -1); - auto init_decoded_values = make_shared(minus_one_value, dense_shape); - auto decoded_values_dense = make_shared(init_decoded_values, decoded_indices, decoded_values); + // compute target labels in a format accepted by OpenVINO CTCLoss + auto logits_shape = make_shared(logits, element::i64); + auto slice_start = make_shared(element::i64, Shape{1}, 0); + auto slice_end = make_shared(element::i64, Shape{1}, 2); + auto slice_step = make_shared(element::i64, Shape{1}, 1); + auto dense_shape = make_shared(logits_shape, slice_start, slice_end, slice_step); + auto minus_one = make_shared(element::i64, Shape{}, -1); + auto labels = make_shared(minus_one, dense_shape)->output(0); + labels = make_shared(labels, decoded_indices, decoded_values); - // Compute label_lenght for each batch - auto minus_one_mask = make_shared(decoded_values_dense, minus_one_value); - auto mask01 = make_shared