From 3d5fe8d44694b6bd0a5b40d1ff051dc139cc60e7 Mon Sep 17 00:00:00 2001 From: Mustafa Cavus Date: Thu, 19 Oct 2023 10:21:28 -0700 Subject: [PATCH] Llm and sd additional ops (#20435) * TorchFX: New ops added (baddbbmm, leaky_relu_) * TorchFX: Initial scaled_dot_product_flash_attention * Code Formatting: scaled_fot_product_attention translation * TorchFX unit test enabled for SDPA * Typo fix in comment line Co-authored-by: Maxim Vafin --------- Co-authored-by: Maxim Vafin --- .../pytorch/torchdynamo/op_support.py | 3 ++ .../src/op/scaled_dot_product_attention.cpp | 36 +++++++++++++++---- src/frontends/pytorch/src/op_table.cpp | 4 +++ .../test_scaled_dot_product_attention.py | 1 + 4 files changed, 37 insertions(+), 7 deletions(-) diff --git a/src/bindings/python/src/openvino/frontend/pytorch/torchdynamo/op_support.py b/src/bindings/python/src/openvino/frontend/pytorch/torchdynamo/op_support.py index 726f3b598bc..4a76d90b160 100644 --- a/src/bindings/python/src/openvino/frontend/pytorch/torchdynamo/op_support.py +++ b/src/bindings/python/src/openvino/frontend/pytorch/torchdynamo/op_support.py @@ -41,6 +41,7 @@ class OperatorSupport(OperatorSupport): "torch.ops.aten.arange.default": None, "torch.ops.aten.argmax.default": None, "torch.ops.aten.avg_pool2d.default": None, + "torch.ops.aten.baddbmm.default": None, "torch.ops.aten.bitwise_and.Tensor": None, "torch.ops.aten.bmm.default": None, "torch.ops.aten.cat.default": None, @@ -67,6 +68,7 @@ class OperatorSupport(OperatorSupport): "torch.ops.aten.hardswish_.default": None, "torch.ops.aten.hardtanh_.default": None, "torch.ops.aten.index.Tensor": None, + "torch.ops.aten.leaky_relu_.default": None, "torch.ops.aten.lift_fresh_copy.default": None, "torch.ops.aten.linalg_vector_norm.default": None, "torch.ops.aten.lt.Tensor": None, @@ -89,6 +91,7 @@ class OperatorSupport(OperatorSupport): "torch.ops.aten.relu.default": None, "torch.ops.aten.relu_.default": None, "torch.ops.aten.rsub.Scalar": None, + "torch.ops.aten._scaled_dot_product_flash_attention.default": None, "torch.ops.aten.select.int": None, "torch.ops.aten.sigmoid.default": None, "torch.ops.aten.silu.default": None, diff --git a/src/frontends/pytorch/src/op/scaled_dot_product_attention.cpp b/src/frontends/pytorch/src/op/scaled_dot_product_attention.cpp index 735324405d1..82231472e40 100644 --- a/src/frontends/pytorch/src/op/scaled_dot_product_attention.cpp +++ b/src/frontends/pytorch/src/op/scaled_dot_product_attention.cpp @@ -15,6 +15,7 @@ #include "openvino/op/matmul.hpp" #include "openvino/op/multiply.hpp" #include "openvino/op/range.hpp" +#include "openvino/op/reshape.hpp" #include "openvino/op/select.hpp" #include "openvino/op/shape_of.hpp" #include "openvino/op/softmax.hpp" @@ -22,6 +23,7 @@ #include "openvino/op/squeeze.hpp" #include "openvino/op/transpose.hpp" #include "openvino/op/unsqueeze.hpp" +#include "openvino/op/util/framework_node.hpp" #include "utils.hpp" namespace ov { @@ -31,10 +33,7 @@ namespace op { using namespace ov::op; -OutputVector translate_scaled_dot_product_attention(const NodeContext& context) { - // aten::scaled_dot_product_attention(Tensor query, Tensor key, Tensor value, Tensor? attn_mask=None, float - // dropout_p=0., bool is_causal=False) - num_inputs_check(context, 6, 6); +std::shared_ptr translate_scaled_dot_product_attention_common(const NodeContext& context) { auto query = context.get_input(0); auto key = context.get_input(1); auto value = context.get_input(2); @@ -68,7 +67,10 @@ OutputVector translate_scaled_dot_product_attention(const NodeContext& context) minus_inf = context.mark_node(std::make_shared(minus_inf, scaled_atten)); // two types of masks are supported. A boolean mask where a value of True indicates that the element should take // part in attention. A float mask of the same type as query, key, value that is added to the attention score. - auto is_causal = context.const_input(5); + auto is_causal = false; + if (!context.input_is_none(5)) { + is_causal = context.const_input(5); + } if (is_causal || !context.input_is_none(3)) { Output mask; Output atten_mask; @@ -100,10 +102,30 @@ OutputVector translate_scaled_dot_product_attention(const NodeContext& context) scaled_atten = context.mark_node(std::make_shared(scaled_atten, atten_mask)); } scaled_atten = context.mark_node(std::make_shared(scaled_atten, -1)); - return {context.mark_node(std::make_shared(scaled_atten, value))}; + return context.mark_node(std::make_shared(scaled_atten, value)); +}; + +OutputVector translate_scaled_dot_product_attention(const NodeContext& context) { + // aten::scaled_dot_product_attention(Tensor query, Tensor key, Tensor value, Tensor? attn_mask=None, float + // dropout_p=0., bool is_causal=False) + num_inputs_check(context, 6, 6); + return {translate_scaled_dot_product_attention_common(context)}; +}; + +OutputVector translate_scaled_dot_product_attention_fx(const NodeContext& context) { + // aten::scaled_dot_product_attention(Tensor query, Tensor key, Tensor value, Tensor? attn_mask=None, float + // dropout_p=0., bool is_causal=False) + num_inputs_check(context, 3, 6); + auto output = translate_scaled_dot_product_attention_common(context); + // TODO: scaled_dot_product_flash_attention has 9 outputs but for most cases only + // the first input is used. Rest of the outputs should be returned properly as + // needed. + ov::OutputVector out_vec; + out_vec.push_back(output); + return {context.mark_node(make_list_construct(out_vec))}; }; } // namespace op } // namespace pytorch } // namespace frontend -} // namespace ov \ No newline at end of file +} // namespace ov diff --git a/src/frontends/pytorch/src/op_table.cpp b/src/frontends/pytorch/src/op_table.cpp index 75665ffe8d4..5614a3881c3 100644 --- a/src/frontends/pytorch/src/op_table.cpp +++ b/src/frontends/pytorch/src/op_table.cpp @@ -213,6 +213,7 @@ OP_CONVERTER(translate_group_norm_fx); OP_CONVERTER(translate_index_fx); OP_CONVERTER(translate_layer_norm_fx); OP_CONVERTER(translate_max_poolnd_fx); +OP_CONVERTER(translate_scaled_dot_product_attention_fx); OP_CONVERTER(translate_slice_fx); OP_CONVERTER(translate_softmax_fx); OP_CONVERTER(translate_transpose_fx); @@ -555,6 +556,7 @@ const std::map get_supported_ops_fx() { {"aten.arange.default", op::translate_arange_fx}, {"aten.argmax.default", op::translate_argmax}, {"aten.avg_pool2d.default", op::translate_avg_poolnd}, + {"aten.baddbmm.default", op::translate_addmm}, {"aten.bitwise_and.Tensor", op::translate_bitwise_and}, {"aten.bmm.default", op::translate_1to1_match_2_inputs_align_types}, {"aten.cat.default", op::translate_cat_fx}, @@ -581,6 +583,7 @@ const std::map get_supported_ops_fx() { {"aten.hardswish_.default", op::inplace_op>}, {"aten.hardtanh_.default", op::inplace_op}, {"aten.index.Tensor", op::translate_index_fx}, + {"aten.leaky_relu_.default", op::inplace_op>}, {"aten.lift_fresh_copy.default", op::skip_node}, {"aten.linalg_vector_norm.default", op::translate_linalg_vector_norm}, {"aten.log.default", op::translate_log}, @@ -603,6 +606,7 @@ const std::map get_supported_ops_fx() { {"aten.relu.default", op::translate_1to1_match_1_inputs}, {"aten.relu_.default", op::inplace_op>}, {"aten.rsub.Scalar", op::translate_rsub}, + {"aten._scaled_dot_product_flash_attention.default", op::translate_scaled_dot_product_attention_fx}, {"aten.select.int", op::translate_select}, {"aten.sigmoid.default", op::translate_1to1_match_1_inputs}, {"aten.silu.default", op::translate_1to1_match_1_inputs}, diff --git a/tests/layer_tests/pytorch_tests/test_scaled_dot_product_attention.py b/tests/layer_tests/pytorch_tests/test_scaled_dot_product_attention.py index 22ed3254718..69c600a0b75 100644 --- a/tests/layer_tests/pytorch_tests/test_scaled_dot_product_attention.py +++ b/tests/layer_tests/pytorch_tests/test_scaled_dot_product_attention.py @@ -36,6 +36,7 @@ class TestScaledDotProductAttention(PytorchLayerTest): @pytest.mark.nightly @pytest.mark.precommit + @pytest.mark.precommit_fx_backend @pytest.mark.parametrize(['mask', "is_causal"], [(False, False), (False, True), (True, True), (True, False)]) def test_scaled_dot_product_atten(self, ie_device, precision, ir_version, mask, is_causal): self._test(*self.create_model(mask, is_causal),ie_device, precision, ir_version)