diff --git a/src/plugins/intel_gpu/include/intel_gpu/op/gemm.hpp b/src/plugins/intel_gpu/include/intel_gpu/op/gemm.hpp index 0f1e6904831..41b610d0785 100644 --- a/src/plugins/intel_gpu/include/intel_gpu/op/gemm.hpp +++ b/src/plugins/intel_gpu/include/intel_gpu/op/gemm.hpp @@ -26,27 +26,12 @@ public: const std::vector& order_c, const ov::element::Type output_type = ov::element::undefined); - Gemm(const ov::Output& A, - const ov::Output& B, - const std::vector& target_shape_a, - const std::vector& target_shape_b, - const std::vector& output_pattern_a, - const std::vector& output_pattern_b, - const std::vector& order_a, - const std::vector& order_b, - const std::vector& order_c, - const ov::element::Type output_type = ov::element::undefined); - bool visit_attributes(ov::AttributeVisitor &visitor) override; void validate_and_infer_types() override; std::shared_ptr clone_with_new_inputs(const ov::OutputVector& new_args) const override; - std::vector get_input0_broadcast_target_shape() const { return m_target_shape_a; } - std::vector get_input1_broadcast_target_shape() const { return m_target_shape_b; } - std::vector get_input0_reshape_pattern() const { return m_output_pattern_a; } - std::vector get_input1_reshape_pattern() const { return m_output_pattern_b; } std::vector get_input0_transpose_order() const { return m_order_a; } std::vector get_input1_transpose_order() const { return m_order_b; } std::vector get_output_transpose_order() const { return m_order_c; } @@ -59,10 +44,6 @@ public: } protected: - std::vector m_target_shape_a; - std::vector m_target_shape_b; - std::vector m_output_pattern_a; - std::vector m_output_pattern_b; std::vector m_order_a; std::vector m_order_b; std::vector m_order_c; @@ -71,10 +52,6 @@ protected: std::vector shape_infer(const Gemm* op, std::vector input_shapes, - const std::vector& target_shape_a, - const std::vector& target_shape_b, - const std::vector& output_pattern_a, - const std::vector& output_pattern_b, const std::vector& order_a, const std::vector& order_b, const std::vector& order_c); diff --git a/src/plugins/intel_gpu/include/intel_gpu/primitives/gemm.hpp b/src/plugins/intel_gpu/include/intel_gpu/primitives/gemm.hpp index 15dd92cd23f..b5d2dd66508 100644 --- a/src/plugins/intel_gpu/include/intel_gpu/primitives/gemm.hpp +++ b/src/plugins/intel_gpu/include/intel_gpu/primitives/gemm.hpp @@ -54,10 +54,6 @@ struct gemm : public primitive_base { : primitive_base(id, inputs, {output_padding}, {optional_data_type{ data_type }}), transpose_input0(transpose_input0 ? 1 : 0), transpose_input1(transpose_input1 ? 1 : 0), - input0_broadcast_target_shape({}), - input1_broadcast_target_shape({}), - input0_reshape_pattern({}), - input1_reshape_pattern({}), alpha(alpha), beta(beta), input_rank(input_rank), @@ -90,10 +86,6 @@ struct gemm : public primitive_base { gemm(const primitive_id& id, const std::vector& inputs, const data_types data_type, - const std::vector& input0_broadcast_target_shape = {}, - const std::vector& input1_broadcast_target_shape = {}, - const std::vector& input0_reshape_pattern = {}, - const std::vector& input1_reshape_pattern = {}, const std::vector& input0_transpose_order = {0, 1, 2, 3}, const std::vector& input1_transpose_order = {0, 1, 2, 3}, const std::vector& output_transpose_order = {}, @@ -101,10 +93,6 @@ struct gemm : public primitive_base { const float beta = 0.0f, const padding& output_padding = padding()) : primitive_base(id, inputs, {output_padding}, {optional_data_type{ data_type }}), - input0_broadcast_target_shape(input0_broadcast_target_shape), - input1_broadcast_target_shape(input1_broadcast_target_shape), - input0_reshape_pattern(input0_reshape_pattern), - input1_reshape_pattern(input1_reshape_pattern), input0_transpose_order(input0_transpose_order), input1_transpose_order(input1_transpose_order), output_transpose_order(output_transpose_order), @@ -133,10 +121,6 @@ struct gemm : public primitive_base { const float beta = 0.0f, const padding& output_padding = padding()) : primitive_base(id, inputs, {output_padding}, {optional_data_type{ data_type }}), - input0_broadcast_target_shape({}), - input1_broadcast_target_shape({}), - input0_reshape_pattern({}), - input1_reshape_pattern({}), input0_transpose_order(input0_transpose_order), input1_transpose_order(input1_transpose_order), output_transpose_order(output_transpose_order), @@ -159,14 +143,6 @@ struct gemm : public primitive_base { uint32_t transpose_input0 = 0; /// @brief Flag for transposing second input matrix uint32_t transpose_input1 = 0; - /// @brief broadcasted target shape of input 0 - std::vector input0_broadcast_target_shape; - /// @brief broadcasted target shape of input 1 - std::vector input1_broadcast_target_shape; - /// @brief reshaped output pattern of input 0 - std::vector input0_reshape_pattern; - /// @brief reshaped output pattern of input 1 - std::vector input1_reshape_pattern; /// @brief order of input 0 std::vector input0_transpose_order; /// @brief order of input 1 @@ -193,10 +169,6 @@ struct gemm : public primitive_base { seed = hash_combine(seed, transpose_input1); seed = hash_combine(seed, indirect_a); seed = hash_combine(seed, indirect_b); - seed = hash_range(seed, input0_broadcast_target_shape.begin(), input0_broadcast_target_shape.end()); - seed = hash_range(seed, input1_broadcast_target_shape.begin(), input1_broadcast_target_shape.end()); - seed = hash_range(seed, input0_reshape_pattern.begin(), input0_reshape_pattern.end()); - seed = hash_range(seed, input1_reshape_pattern.begin(), input1_reshape_pattern.end()); seed = hash_range(seed, input0_transpose_order.begin(), input0_transpose_order.end()); seed = hash_range(seed, input1_transpose_order.begin(), input1_transpose_order.end()); seed = hash_range(seed, output_transpose_order.begin(), output_transpose_order.end()); @@ -225,10 +197,6 @@ struct gemm : public primitive_base { primitive_base::save(ob); ob << transpose_input0; ob << transpose_input1; - ob << input0_broadcast_target_shape; - ob << input1_broadcast_target_shape; - ob << input0_reshape_pattern; - ob << input1_reshape_pattern; ob << input0_transpose_order; ob << input1_transpose_order; ob << output_transpose_order; @@ -246,10 +214,6 @@ struct gemm : public primitive_base { primitive_base::load(ib); ib >> transpose_input0; ib >> transpose_input1; - ib >> input0_broadcast_target_shape; - ib >> input1_broadcast_target_shape; - ib >> input0_reshape_pattern; - ib >> input1_reshape_pattern; ib >> input0_transpose_order; ib >> input1_transpose_order; ib >> output_transpose_order; diff --git a/src/plugins/intel_gpu/src/graph/gemm.cpp b/src/plugins/intel_gpu/src/graph/gemm.cpp index 49f0fefd0f1..bd19dcf75b7 100644 --- a/src/plugins/intel_gpu/src/graph/gemm.cpp +++ b/src/plugins/intel_gpu/src/graph/gemm.cpp @@ -10,18 +10,6 @@ #include "intel_gpu/op/gemm.hpp" -namespace { -template ::value>::type> -int find_index_from_vec(const std::vector& vec, const DT value) { - int idx = 0; - for (auto v : vec) { - if (v != static_cast(value)) - break; - idx += 1; - } - return idx; -} -} // namespace namespace cldnn { GPU_DEFINE_PRIMITIVE_TYPE_ID(gemm) @@ -139,10 +127,6 @@ std::vector gemm_inst::calc_output_layouts(gemm_node const& node, const std::vector output_shapes = ov::intel_gpu::op::shape_infer(&op, input_shapes, - prim->input0_broadcast_target_shape, - prim->input1_broadcast_target_shape, - prim->input0_reshape_pattern, - prim->input1_reshape_pattern, prim->input0_transpose_order, prim->input1_transpose_order, prim->output_transpose_order); @@ -158,28 +142,6 @@ template std::vector gemm_inst::calc_output_layouts(ge std::vector gemm_inst::transform_input_layouts(const std::shared_ptr primitive, const std::vector& input_layouts) { - auto get_reshaped_input_shape = [&](const ov::PartialShape& input_pshape, - const std::vector& broadcast_target_shape, - const std::vector& reshape_pattern) { - ov::PartialShape reshaped_input_pshape; - - if (broadcast_target_shape.size() > 0 && reshape_pattern.size() > 0) { - std::vector dims(input_pshape); - int idx_recalc = find_index_from_vec(broadcast_target_shape, 1); - int idx_target = find_index_from_vec(reshape_pattern, 0); - if (dims[idx_recalc].is_static() && dims[idx_target].is_static()) { - dims[idx_recalc] *= dims[idx_target]; - } else { - dims[idx_recalc] = ov::Dimension::dynamic(); - } - dims.erase(dims.begin() + idx_target); - reshaped_input_pshape = ov::PartialShape(dims); - } else { - reshaped_input_pshape = input_pshape; - } - return reshaped_input_pshape; - }; - auto get_transposed_input_shape = [&](const ov::PartialShape& input_pshape, size_t input_rank, size_t output_rank, bool transpose, bool first_input) { ov::PartialShape transposed_input_pshape; @@ -214,30 +176,20 @@ std::vector gemm_inst::transform_input_layouts(const std::shared_ptrinput0_broadcast_target_shape, - primitive->input0_reshape_pattern); - auto reshaped_input1_pshape = get_reshaped_input_shape(input_layouts[1].get_partial_shape(), - primitive->input1_broadcast_target_shape, - primitive->input1_reshape_pattern); + auto input0_pshape = input_layouts[0].get_partial_shape(); + auto input1_pshape = input_layouts[1].get_partial_shape(); bool reordered = primitive->input_rank > 4 || primitive->weight_rank > 4; size_t output_rank = std::max(primitive->input_rank, primitive->weight_rank); size_t input_rank = reordered ? output_rank : primitive->input_rank; size_t weight_rank = reordered ? output_rank : primitive->weight_rank; - auto transposed_input0_pshape = get_transposed_input_shape(reshaped_input0_pshape, input_rank, output_rank, primitive->transpose_input0, true); - auto transposed_input1_pshape = get_transposed_input_shape(reshaped_input1_pshape, weight_rank, output_rank, primitive->transpose_input1, false); + auto transposed_input0_pshape = get_transposed_input_shape(input0_pshape, input_rank, output_rank, primitive->transpose_input0, true); + auto transposed_input1_pshape = get_transposed_input_shape(input1_pshape, weight_rank, output_rank, primitive->transpose_input1, false); std::vector layouts = input_layouts; layouts[0].set_partial_shape(transposed_input0_pshape); - if (primitive->input0_broadcast_target_shape.size() > input_rank) { - layouts[0].format = format::adjust_to_rank(layouts[0].format, input_rank); - } layouts[1].set_partial_shape(transposed_input1_pshape); - if (primitive->input1_broadcast_target_shape.size() > weight_rank) { - layouts[1].format = format::adjust_to_rank(layouts[1].format, weight_rank); - } if (primitive->input_size() == 3) { auto bias_pshape = input_layouts[2].get_partial_shape(); diff --git a/src/plugins/intel_gpu/src/graph/impls/ocl/gemm.cpp b/src/plugins/intel_gpu/src/graph/impls/ocl/gemm.cpp index 03124262072..8d0efede6ee 100644 --- a/src/plugins/intel_gpu/src/graph/impls/ocl/gemm.cpp +++ b/src/plugins/intel_gpu/src/graph/impls/ocl/gemm.cpp @@ -2,6 +2,8 @@ // SPDX-License-Identifier: Apache-2.0 // +#include "intel_gpu/op/gemm.hpp" +#include "intel_gpu/plugin/common_utils.hpp" #include "intel_gpu/graph/kernel_impl_params.hpp" #include "multi_stage_primitive.hpp" @@ -173,14 +175,46 @@ public: params.beta = primitive->beta; params.transpose_input0 = primitive->transpose_input0; params.transpose_input1 = primitive->transpose_input1; - params.input0_target_shape = primitive->input0_broadcast_target_shape; - params.input1_target_shape = primitive->input1_broadcast_target_shape; - params.input0_output_pattern = primitive->input0_reshape_pattern; - params.input1_output_pattern = primitive->input0_reshape_pattern; params.input0_order = primitive->input0_transpose_order; params.input1_order = primitive->input1_transpose_order; params.output_order = primitive->output_transpose_order; + auto input0_pshape = impl_param.input_layouts[0].get_partial_shape(); + auto input1_pshape = impl_param.input_layouts[1].get_partial_shape(); + const auto is_broadcastable = input0_pshape.rank().is_static() && + input1_pshape.rank().is_static() && + input0_pshape.size() > 1 && + input1_pshape.size() > 1 && + (primitive->input_rank == primitive->weight_rank); + if (is_broadcastable) { + auto transpose_pshape = [](const ov::PartialShape pshape, const std::vector& order) { + auto transposed_pshape = ov::PartialShape::dynamic(pshape.rank()); + for (size_t i = 0; i < order.size(); i++) { + transposed_pshape[i] = pshape[order[i]]; + } + return transposed_pshape; + }; + size_t max_rank = input0_pshape.size(); + auto default_order = ov::intel_gpu::op::Gemm::default_order(max_rank); + auto input0_trans_pshape = (primitive->input0_transpose_order != default_order) ? + transpose_pshape(input0_pshape, primitive->input0_transpose_order) : + input0_pshape; + auto input1_trans_pshape = (primitive->input1_transpose_order != default_order) ? + transpose_pshape(input1_pshape, primitive->input1_transpose_order) : + input1_pshape; + for (size_t i = 0; i < max_rank - 2; ++i) { + if (input0_trans_pshape[i].is_static() && input1_trans_pshape[i].is_static()) { + if (input1_trans_pshape[i].get_length() > input0_trans_pshape[i].get_length()) { + params.input0_reshape_axes = primitive->input0_transpose_order[i]; + params.input0_broadcast_val = input1_trans_pshape[i].get_length() / input0_trans_pshape[i].get_length(); + } else if (input0_trans_pshape[i].get_length() > input1_trans_pshape[i].get_length()) { + params.input1_reshape_axes = primitive->input1_transpose_order[i]; + params.input1_broadcast_val = input0_trans_pshape[i].get_length() / input1_trans_pshape[i].get_length(); + } + } + } + } + params.indirect_input0 = primitive->indirect_a && indirect; params.indirect_input1 = primitive->indirect_b && indirect; if (indirect && (primitive->indirect_a || primitive->indirect_b)) { diff --git a/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.cpp b/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.cpp index 44eba6cfbc5..461c306a567 100644 --- a/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.cpp +++ b/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.cpp @@ -215,41 +215,37 @@ JitConstants GemmKernelBase::GetJitConstants(const gemm_params& params) const { jit.AddConstant(MakeJitConstant("BIAS_TERM", 1)); } - auto get_broadcast_input_str = [](const std::vector& target_shape) { - const size_t target_rank = target_shape.size(); + auto get_broadcast_input_str = [](const size_t input_rank, const int64_t axes, const int64_t val) { std::vector dims; - if (target_rank == 1) { + if (input_rank == 1) { dims = {"x"}; - } else if (target_rank == 2) { + } else if (input_rank == 2) { dims = {"y", "x"}; - } else if (target_rank == 3) { + } else if (input_rank == 3) { dims = {"f", "y", "x"}; - } else if (target_rank == 4) { + } else if (input_rank == 4) { dims = {"b", "f", "y", "x"}; - } else if (target_rank == 5) { + } else if (input_rank == 5) { dims = {"b", "f", "z", "y", "x"}; - } else if (target_rank == 6) { + } else if (input_rank == 6) { dims = {"b", "f", "w", "z", "y", "x"}; } - int pos = 0; - for (auto ts : target_shape) { - if (ts != 1) - break; - pos += 1; - } - std::string str = dims[pos] + " /= " + std::to_string(target_shape[pos]) + ";"; - return str; + return dims[axes] + " /= " + std::to_string(val) + ";"; }; - if (params.input0_target_shape.size() > 1) { + if (params.input0_broadcast_val != 0) { jit.AddConstants({ MakeJitConstant("BROADCAST_INPUT0", true), - MakeJitConstant("DO_BROADCAST_INPUT0", get_broadcast_input_str(params.input0_target_shape)), + MakeJitConstant("DO_BROADCAST_INPUT0", get_broadcast_input_str(params.inputs[0].GetDims().size(), + params.input0_reshape_axes, + params.input0_broadcast_val)), }); } - if (params.input1_target_shape.size() > 1) { + if (params.input1_broadcast_val != 0) { jit.AddConstants({ MakeJitConstant("BROADCAST_INPUT1", true), - MakeJitConstant("DO_BROADCAST_INPUT1", get_broadcast_input_str(params.input1_target_shape)), + MakeJitConstant("DO_BROADCAST_INPUT1", get_broadcast_input_str(params.inputs[1].GetDims().size(), + params.input1_reshape_axes, + params.input1_broadcast_val)), }); } diff --git a/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.h b/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.h index 633c8171c99..32a186412d3 100644 --- a/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.h +++ b/src/plugins/intel_gpu/src/kernel_selector/kernels/gemm/gemm_kernel_base.h @@ -19,13 +19,13 @@ struct gemm_params : public base_params { float beta; uint32_t transpose_input0; uint32_t transpose_input1; - std::vector input0_target_shape; - std::vector input1_target_shape; - std::vector input0_output_pattern; - std::vector input1_output_pattern; std::vector input0_order; std::vector input1_order; std::vector output_order; + int64_t input0_reshape_axes = 0; + int64_t input1_reshape_axes = 0; + int64_t input0_broadcast_val = 0; + int64_t input1_broadcast_val = 0; DataTensor beam_table; bool indirect_input0 = false; bool indirect_input1 = false; diff --git a/src/plugins/intel_gpu/src/plugin/ops/matmul.cpp b/src/plugins/intel_gpu/src/plugin/ops/matmul.cpp index d455c1fa839..05335d2193b 100644 --- a/src/plugins/intel_gpu/src/plugin/ops/matmul.cpp +++ b/src/plugins/intel_gpu/src/plugin/ops/matmul.cpp @@ -156,23 +156,17 @@ static void CreateGemmOp(ProgramBuilder& p, const std::shared_ptrget_input_partial_shape(1); auto out_shape = op->get_output_partial_shape(0); - size_t rank_a = op->get_input0_reshape_pattern().size() > 0 ? op->get_input0_reshape_pattern().size() - : shape_a.rank().get_length(); - size_t rank_b = op->get_input1_reshape_pattern().size() > 0 ? op->get_input1_reshape_pattern().size() - :shape_b.rank().get_length(); + size_t rank_a = shape_a.rank().get_length(); + size_t rank_b = shape_b.rank().get_length(); size_t output_rank = out_shape.rank().get_length(); - OPENVINO_ASSERT(rank_a == op->get_input0_transpose_order().size(), "[GPU] Length of input0_order is not same as rank of input0"); - OPENVINO_ASSERT(rank_b == op->get_input1_transpose_order().size(), "[GPU] Length of input1_order is not same as rank of input1"); - OPENVINO_ASSERT(output_rank == op->get_output_transpose_order().size(), "[GPU] Length of output_order is not same as rank of output"); + OPENVINO_ASSERT(rank_a == op->get_input0_transpose_order().size(), "[GPU] Length of input0_transpose_order is not same as rank of input0"); + OPENVINO_ASSERT(rank_b == op->get_input1_transpose_order().size(), "[GPU] Length of input1_transpose_order is not same as rank of input1"); + OPENVINO_ASSERT(output_rank == op->get_output_transpose_order().size(), "[GPU] Length of output_transpose_order is not same as rank of output"); auto gemmPrim = cldnn::gemm(layerName, inputs, cldnn::element_type_to_data_type(op->get_output_element_type(0)), - op->get_input0_broadcast_target_shape(), - op->get_input1_broadcast_target_shape(), - op->get_input0_reshape_pattern(), - op->get_input1_reshape_pattern(), op->get_input0_transpose_order(), op->get_input1_transpose_order(), op->get_output_transpose_order(), diff --git a/src/plugins/intel_gpu/src/plugin/transformations/op/gemm.cpp b/src/plugins/intel_gpu/src/plugin/transformations/op/gemm.cpp index 16803da2ef3..7c0622ec0e9 100644 --- a/src/plugins/intel_gpu/src/plugin/transformations/op/gemm.cpp +++ b/src/plugins/intel_gpu/src/plugin/transformations/op/gemm.cpp @@ -3,6 +3,7 @@ // #include "intel_gpu/op/gemm.hpp" +#include "intel_gpu/plugin/common_utils.hpp" #include "matmul_shape_inference.hpp" #include "broadcast_shape_inference.hpp" #include "reshape_shape_inference.hpp" @@ -22,35 +23,6 @@ Gemm::Gemm(const ov::Output& A, const std::vector& order_c, const ov::element::Type output_type) : ov::op::v0::MatMul() - , m_target_shape_a({}) - , m_target_shape_b({}) - , m_output_pattern_a({}) - , m_output_pattern_b({}) - , m_order_a(order_a) - , m_order_b(order_b) - , m_order_c(order_c) - , m_output_type(output_type) { - set_arguments({A, B}); - set_transpose_a(false); - set_transpose_b(false); - validate_and_infer_types(); -} - -Gemm::Gemm(const ov::Output& A, - const ov::Output& B, - const std::vector& target_shape_a, - const std::vector& target_shape_b, - const std::vector& output_pattern_a, - const std::vector& output_pattern_b, - const std::vector& order_a, - const std::vector& order_b, - const std::vector& order_c, - const ov::element::Type output_type) - : ov::op::v0::MatMul() - , m_target_shape_a(target_shape_a) - , m_target_shape_b(target_shape_b) - , m_output_pattern_a(output_pattern_a) - , m_output_pattern_b(output_pattern_b) , m_order_a(order_a) , m_order_b(order_b) , m_order_c(order_c) @@ -64,16 +36,7 @@ Gemm::Gemm(const ov::Output& A, std::shared_ptr Gemm::clone_with_new_inputs(const ov::OutputVector& new_args) const { check_new_args_count(this, new_args); - return std::make_shared(new_args.at(0), - new_args.at(1), - m_target_shape_a, - m_target_shape_b, - m_output_pattern_a, - m_output_pattern_b, - m_order_a, - m_order_b, - m_order_c, - m_output_type); + return std::make_shared(new_args.at(0), new_args.at(1), m_order_a, m_order_b, m_order_c, m_output_type); } void Gemm::validate_and_infer_types() { @@ -86,10 +49,6 @@ void Gemm::validate_and_infer_types() { auto out_shapes = shape_infer(this, std::vector{get_input_partial_shape(0), get_input_partial_shape(1)}, - m_target_shape_a, - m_target_shape_b, - m_output_pattern_a, - m_output_pattern_b, m_order_a, m_order_b, m_order_c); @@ -108,60 +67,47 @@ bool Gemm::visit_attributes(ov::AttributeVisitor &visitor) { std::vector shape_infer(const Gemm* op, std::vector input_shapes, - const std::vector& target_shape_a, - const std::vector& target_shape_b, - const std::vector& output_pattern_a, - const std::vector& output_pattern_b, const std::vector& order_a, const std::vector& order_b, const std::vector& order_c) { auto shape_a = input_shapes[0]; auto shape_b = input_shapes[1]; - // broadcasted shapes - auto broadcast_shape = [](const ov::PartialShape shape, const std::vector& target_shape) { - ov::op::v3::Broadcast broadcast; - auto tshape = target_shape; - broadcast.set_broadcast_spec(ov::op::BroadcastType::BIDIRECTIONAL); - std::unordered_map const_data; - const_data.emplace(1, ov::Tensor(ov::element::i32, ov::Shape{tshape.size()}, static_cast(tshape.data()))); - return ov::op::v3::shape_infer(&broadcast, - std::vector{shape, ov::PartialShape(ov::Shape{tshape.size()})}, - ov::make_tensor_accessor(const_data)); - }; - auto shape_a_b = (target_shape_a.size() > 1) ? broadcast_shape(shape_a, target_shape_a)[0] : shape_a; - auto shape_b_b = (target_shape_b.size() > 1) ? broadcast_shape(shape_b, target_shape_b)[0] : shape_b; - - // reshaped shapes - auto reshape_shape = [](const ov::PartialShape shape, const std::vector& output_pattern) { - ov::op::v1::Reshape reshape; - auto opattern = output_pattern; - reshape.set_special_zero(true); - std::unordered_map const_data; - const_data.emplace(1, ov::Tensor(ov::element::i64, ov::Shape{opattern.size()}, static_cast(opattern.data()))); - return ov::op::v1::shape_infer(&reshape, - std::vector{shape, ov::PartialShape(ov::Shape{opattern.size()})}, - ov::make_tensor_accessor(const_data)); - }; - auto shape_a_r = (output_pattern_a.size() > 1) ? reshape_shape(shape_a_b, output_pattern_a)[0] : shape_a_b; - auto shape_b_r = (output_pattern_b.size() > 1) ? reshape_shape(shape_b_b, output_pattern_b)[0] : shape_b_b; - - // transposed shapes - auto transpose_shape = [](const ov::PartialShape shape, const std::vector& order) { - auto shape_transposed = ov::PartialShape::dynamic(shape.rank()); + // transposed shape + auto transpose_pshape = [](const ov::PartialShape pshape, const std::vector& order) { + auto transposed_pshape = ov::PartialShape::dynamic(pshape.rank()); for (size_t i = 0; i < order.size(); i++) { - shape_transposed[i] = shape[order[i]]; + transposed_pshape[i] = pshape[order[i]]; } - return shape_transposed; + return transposed_pshape; }; - auto shape_a_t = (order_a.size() > 1) ? transpose_shape(shape_a_r, order_a) : shape_a_r; - auto shape_b_t = (order_b.size() > 1) ? transpose_shape(shape_b_r, order_b) : shape_b_r; + + auto shape_a_t = (order_a.size() > 1) ? transpose_pshape(shape_a, order_a) : shape_a; + auto shape_b_t = (order_b.size() > 1) ? transpose_pshape(shape_b, order_b) : shape_b; + + // broadcast all batch dimensions + const auto is_broadcastable = shape_a_t.rank().is_static() && + shape_a_t.rank().is_static() && + shape_a_t.size() > 1 && + shape_b_t.size() > 1 && + (shape_a_t.size() == shape_b_t.size()); + if (is_broadcastable) { + size_t max_rank = shape_a_t.size(); + for (size_t i = 0; i < max_rank - 2; ++i) { + if (shape_a_t[i].is_static() && shape_b_t[i].is_static()) { + auto result = std::max(shape_a_t[i].get_length(), shape_b_t[i].get_length()); + shape_a_t[i] = result; + shape_b_t[i] = result; + } + } + } + OPENVINO_ASSERT(op != nullptr, "op should not be nullptr for shape_infer."); auto out_shapes = ov::op::v0::shape_infer(dynamic_cast(op), std::vector{shape_a_t, shape_b_t}); if (order_c.size() > 0) { - return { transpose_shape(out_shapes[0], order_c) }; + return { transpose_pshape(out_shapes[0], order_c) }; } else { return { out_shapes[0] }; } diff --git a/src/plugins/intel_gpu/src/plugin/transformations/op/indirect_gemm.cpp b/src/plugins/intel_gpu/src/plugin/transformations/op/indirect_gemm.cpp index bd557f811e6..80e8a7a602c 100644 --- a/src/plugins/intel_gpu/src/plugin/transformations/op/indirect_gemm.cpp +++ b/src/plugins/intel_gpu/src/plugin/transformations/op/indirect_gemm.cpp @@ -50,10 +50,6 @@ void IndirectGemm::validate_and_infer_types() { auto out_shapes = shape_infer(this, std::vector{get_input_partial_shape(0), get_input_partial_shape(1)}, - m_target_shape_a, - m_target_shape_b, - m_output_pattern_a, - m_output_pattern_b, m_order_a, m_order_b, m_order_c); diff --git a/src/plugins/intel_gpu/src/plugin/transformations/broadcast_reshape_matmul_fusion.cpp b/src/plugins/intel_gpu/src/plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.cpp similarity index 81% rename from src/plugins/intel_gpu/src/plugin/transformations/broadcast_reshape_matmul_fusion.cpp rename to src/plugins/intel_gpu/src/plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.cpp index 17df3d3d1a7..9cbafc2f229 100644 --- a/src/plugins/intel_gpu/src/plugin/transformations/broadcast_reshape_matmul_fusion.cpp +++ b/src/plugins/intel_gpu/src/plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.cpp @@ -2,7 +2,7 @@ // SPDX-License-Identifier: Apache-2.0 // -#include "broadcast_reshape_matmul_fusion.hpp" +#include "unsqueeze_broadcast_reshape_matmul_fusion.hpp" #include "intel_gpu/op/gemm.hpp" @@ -10,6 +10,7 @@ #include "openvino/op/broadcast.hpp" #include "openvino/op/constant.hpp" #include "openvino/op/reshape.hpp" +#include "openvino/op/unsqueeze.hpp" #include "openvino/pass/pattern/op/wrap_type.hpp" #include "openvino/pass/pattern/op/or.hpp" #include "transformations/utils/utils.hpp" @@ -17,7 +18,7 @@ namespace ov { namespace intel_gpu { -BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { +UnsqueezeBroadcastReshapeMatmulFusion::UnsqueezeBroadcastReshapeMatmulFusion() { using namespace ov::pass::pattern; auto not_reshape = [](const ov::Output& output) -> bool { @@ -35,10 +36,15 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { auto input_a_m = any_input(not_reshape); auto input_b_m = any_input(not_reshape); + auto unsqueeze_a_axes_m = wrap_type(); + auto unsqueeze_a_m = wrap_type({input_a_m, unsqueeze_a_axes_m}, consumers_count(1)); + auto unsqueeze_b_axes_m = wrap_type(); + auto unsqueeze_b_m = wrap_type({input_b_m, unsqueeze_b_axes_m}, consumers_count(1)); + auto broadcast_a_target_shape_m = wrap_type(); - auto broadcast_a_m = wrap_type({input_a_m, broadcast_a_target_shape_m}, broadcast_rank_equals_and_has_static_dims); + auto broadcast_a_m = wrap_type({unsqueeze_a_m, broadcast_a_target_shape_m}, broadcast_rank_equals_and_has_static_dims); auto broadcast_b_target_shape_m = wrap_type(); - auto broadcast_b_m = wrap_type({input_b_m, broadcast_b_target_shape_m}, broadcast_rank_equals_and_has_static_dims); + auto broadcast_b_m = wrap_type({unsqueeze_b_m, broadcast_b_target_shape_m}, broadcast_rank_equals_and_has_static_dims); auto reshape_a_pattern_m = wrap_type(); auto reshape_a_m = wrap_type({broadcast_a_m, reshape_a_pattern_m}, reshape_rank_equals_and_has_static_dim); @@ -57,8 +63,6 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { return false; } - auto target_shape_a = std::vector(); - auto target_shape_b = std::vector(); size_t input_a_output_idx = matmul->get_input_source_output(0).get_index(); size_t input_b_output_idx = matmul->get_input_source_output(1).get_index(); auto order_a = matmul->get_input0_transpose_order(); @@ -68,13 +72,31 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { return order.size() == 4 && order[1] == 2; }; + if (pattern_map.count(unsqueeze_a_m) > 0) { + if (!valid_transpose_order(order_a)) + return false; + auto unsqueeze_a = std::dynamic_pointer_cast(pattern_map.at(unsqueeze_a_m).get_node_shared_ptr()); + if (!unsqueeze_a) + return false; + input_a_output_idx = unsqueeze_a->get_input_source_output(0).get_index(); + } + if (pattern_map.count(unsqueeze_b_m) > 0) { + if (!valid_transpose_order(order_b)) + return false; + auto unsqueeze_b = std::dynamic_pointer_cast(pattern_map.at(unsqueeze_b_m).get_node_shared_ptr()); + if (!unsqueeze_b) + return false; + input_b_output_idx = unsqueeze_b->get_input_source_output(0).get_index(); + } + + auto target_shape_a = std::vector(); + auto target_shape_b = std::vector(); + auto valid_broadcast_target_shape = [](const std::vector& target_shape) { return std::count_if(target_shape.begin(), target_shape.end(), [](int32_t s) { return s != 1; }) == 1; }; if (pattern_map.count(broadcast_a_m) > 0) { - if (!valid_transpose_order(order_a)) - return false; auto broadcast_a = std::dynamic_pointer_cast(pattern_map.at(broadcast_a_m).get_node_shared_ptr()); if (!broadcast_a || broadcast_a->get_broadcast_spec().m_type != ov::op::BroadcastType::BIDIRECTIONAL) return false; @@ -82,11 +104,8 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { target_shape_a = broadcast_a_target_shape->cast_vector(); if (!valid_broadcast_target_shape(target_shape_a)) return false; - input_a_output_idx = broadcast_a->get_input_source_output(0).get_index(); } if (pattern_map.count(broadcast_b_m) > 0) { - if (!valid_transpose_order(order_b)) - return false; auto broadcast_b = std::dynamic_pointer_cast(pattern_map.at(broadcast_b_m).get_node_shared_ptr()); if (!broadcast_b || broadcast_b->get_broadcast_spec().m_type != ov::op::BroadcastType::BIDIRECTIONAL) return false; @@ -94,7 +113,6 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { target_shape_b = broadcast_b_target_shape->cast_vector(); if (!valid_broadcast_target_shape(target_shape_b)) return false; - input_b_output_idx = broadcast_b->get_input_source_output(0).get_index(); } auto pattern_a = std::vector(); @@ -123,10 +141,6 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { auto gemm = std::make_shared(input_a, input_b, - target_shape_a, - target_shape_b, - pattern_a, - pattern_b, order_a, order_b, order_c); @@ -137,7 +151,7 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() { return true; }; - auto m = std::make_shared(matmul_m, "BroadcastReshapeMatmulFusion"); + auto m = std::make_shared(matmul_m, "UnsqueezeBroadcastReshapeMatmulFusion"); this->register_matcher(m, callback); } diff --git a/src/plugins/intel_gpu/src/plugin/transformations/broadcast_reshape_matmul_fusion.hpp b/src/plugins/intel_gpu/src/plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.hpp similarity index 56% rename from src/plugins/intel_gpu/src/plugin/transformations/broadcast_reshape_matmul_fusion.hpp rename to src/plugins/intel_gpu/src/plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.hpp index e3ad540e0a4..35ed30cdc97 100644 --- a/src/plugins/intel_gpu/src/plugin/transformations/broadcast_reshape_matmul_fusion.hpp +++ b/src/plugins/intel_gpu/src/plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.hpp @@ -9,10 +9,10 @@ namespace ov { namespace intel_gpu { -class BroadcastReshapeMatmulFusion : public ov::pass::MatcherPass { +class UnsqueezeBroadcastReshapeMatmulFusion : public ov::pass::MatcherPass { public: - OPENVINO_RTTI("BroadcastReshapeMatmulFusion", "0"); - BroadcastReshapeMatmulFusion(); + OPENVINO_RTTI("UnsqueezeBroadcastReshapeMatmulFusion", "0"); + UnsqueezeBroadcastReshapeMatmulFusion(); }; } // namespace intel_gpu diff --git a/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp b/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp index 13826be77cb..94e90b475af 100644 --- a/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp +++ b/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp @@ -60,7 +60,7 @@ #include "plugin/transformations/transpose_matmul_fusion.hpp" #include "plugin/transformations/indirect_kv_cache.hpp" #include "plugin/transformations/convert_convolution.hpp" -#include "plugin/transformations/broadcast_reshape_matmul_fusion.hpp" +#include "plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.hpp" #include "transformations/common_optimizations/broadcast_elementwise_fusion.hpp" #include "transformations/common_optimizations/broadcast_transition.hpp" #include "transformations/common_optimizations/common_optimizations.hpp" @@ -719,14 +719,13 @@ void TransformationsPipeline::apply(std::shared_ptr func) { manager.register_pass(device_info.max_work_group_size); manager.register_pass(); manager.register_pass(); - if (!device_info.supports_immad) + if (!device_info.supports_immad) { manager.register_pass(); + manager.register_pass(); + } manager.register_pass(); - manager.register_pass(); manager.register_pass(); - if (!device_info.supports_immad) - manager.register_pass(); const size_t zp_pad_size = device_info.supports_immad ? 16 : 32; manager.register_pass(zp_pad_size); diff --git a/src/plugins/intel_gpu/tests/unit/test_cases/gemm_gpu_test.cpp b/src/plugins/intel_gpu/tests/unit/test_cases/gemm_gpu_test.cpp index de1f67ce0d6..2a95deed914 100644 --- a/src/plugins/intel_gpu/tests/unit/test_cases/gemm_gpu_test.cpp +++ b/src/plugins/intel_gpu/tests/unit/test_cases/gemm_gpu_test.cpp @@ -847,7 +847,7 @@ public: } } - void test_broadcast_transpose_matmul(bool is_caching_test) { + void test_unsqueeze_broadcast_reshape_transpose_matmul(bool is_caching_test) { tests::random_generator rg; rg.set_seed(GET_SUITE_NAME); @@ -876,21 +876,19 @@ public: auto& engine = get_test_engine(); ov::Shape input0_shape; ov::Shape input1_shape; - std::vector input1_target_shape; std::vector input0_order; std::vector input1_order; ov::Shape beam_table_shape; cldnn::layout input0_layout; cldnn::layout input1_layout; - input0_shape = { BATCH_SIZE, 16, M_SIZE, K_SIZE }; - input1_shape = { N_SIZE, BATCH_SIZE, 1, K_SIZE }; - input1_target_shape = { 1, 1, 16, 1 }; + input0_shape = { BATCH_SIZE, 32, M_SIZE, K_SIZE }; + input1_shape = { N_SIZE, BATCH_SIZE, 2, K_SIZE }; input0_order = { 0, 1, 2, 3 }; input1_order = { 1, 2, 3, 0 }; - input0_layout = layout{ov::PartialShape::dynamic(input0_shape.size()), data_types::f32, format::bfyx}; - input1_layout = layout{ov::PartialShape::dynamic(input1_shape.size()), data_types::f32, format::bfyx}; + input0_layout = layout{ov::PartialShape{ov::Dimension::dynamic(), 32, ov::Dimension::dynamic(), K_SIZE}, data_types::f32, format::bfyx}; + input1_layout = layout{ov::PartialShape{ov::Dimension::dynamic(), ov::Dimension::dynamic(), 2, K_SIZE}, data_types::f32, format::bfyx}; auto input0_mem = engine.allocate_memory(layout{ov::PartialShape(input0_shape), data_types::f32, format::bfyx}); auto input1_mem = engine.allocate_memory(layout{ov::PartialShape(input1_shape), data_types::f32, format::bfyx}); @@ -904,142 +902,7 @@ public: topology topology; topology.add(input_layout("input0", input0_layout), input_layout("input1", input1_layout), - gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f32, {}, input1_target_shape, {}, {}, input0_order, input1_order) - ); - - ExecutionConfig config = get_test_default_config(engine); - config.set_property(ov::intel_gpu::optimize_data(true)); - config.set_property(ov::intel_gpu::allow_new_shape_infer(true)); - network::ptr network = get_network(engine, topology, config, get_test_stream_ptr(), is_caching_test); - network->set_input_data("input0", input0_mem); - network->set_input_data("input1", input1_mem); - - auto inst = network->get_primitive("gemm"); - auto impl = inst->get_impl(); - ASSERT_TRUE(impl != nullptr); - - auto outputs = network->execute(); - - auto output_mem = outputs.at("gemm").get_memory(); - cldnn::mem_lock output_ptr(output_mem, get_test_stream()); - - ov::Shape ref_input0_shape; - ov::Shape ref_input1_broadcasted_shape; - ov::Shape ref_input1_shape; - ov::Shape ref_output_shape; - - ref_input0_shape = { BATCH_SIZE, 16, M_SIZE, K_SIZE }; - ref_input1_broadcasted_shape = { N_SIZE, BATCH_SIZE, 16, K_SIZE }; - ref_input1_shape = { BATCH_SIZE, 16, K_SIZE, N_SIZE }; - ref_output_shape = { BATCH_SIZE, 16, M_SIZE, N_SIZE }; - - std::vector ref_out_data; - ref_out_data.resize(ov::shape_size(ref_output_shape)); - - std::vector ref_input_0_data(input_0_data.size()); - std::vector ref_input_1_broadcasted_data(ov::shape_size(ref_input1_broadcasted_shape)); - std::vector ref_input_1_data(ref_input_1_broadcasted_data.size()); - - ov::reference::transpose((const char *)(input_0_data.data()), - (char *)(ref_input_0_data.data()), - input0_shape, - sizeof(float), - input0_order, - ref_input0_shape); - - ov::reference::broadcast(reinterpret_cast(input_1_data.data()), - reinterpret_cast(ref_input_1_broadcasted_data.data()), - input1_shape, - ref_input1_broadcasted_shape, - ov::AxisSet({}), - sizeof(float)); - - ov::reference::transpose((const char *)(ref_input_1_broadcasted_data.data()), - (char *)(ref_input_1_data.data()), - ref_input1_broadcasted_shape, - sizeof(float), - input1_order, - ref_input1_shape); - - ov::reference::matmul(ref_input_0_data.data(), - ref_input_1_data.data(), - ref_out_data.data(), - ref_input0_shape, - ref_input1_shape, - ref_output_shape, - false, - false); - - ASSERT_EQ(output_ptr.size(), ref_out_data.size()); - - const auto abs_error = 0.0001; - for (uint32_t i = 0; i < ref_out_data.size(); ++i) { - ASSERT_NEAR(output_ptr[i], ref_out_data[i], abs_error) << "at " << i; - } - } - - void test_broadcast_reshape_transpose_matmul(bool is_caching_test) { - tests::random_generator rg; - rg.set_seed(GET_SUITE_NAME); - - const unsigned long BATCH_SIZE = 1; - const unsigned long M_SIZE = 1; - const unsigned long K_SIZE = 32; - const unsigned long N_SIZE = 21; - - auto fill_mem = [&](cldnn::memory_ptr mem, std::vector& data) { - cldnn::mem_lock mem_ptr(mem, get_test_stream()); - auto&& l = mem->get_layout(); - auto data_idx = 0; - for (cldnn::tensor::value_type b = 0; b < l.batch(); ++b) { - for (cldnn::tensor::value_type f = 0; f < l.feature(); ++f) { - for (cldnn::tensor::value_type z = 0; z < l.spatial(2); ++z) { - for (cldnn::tensor::value_type y = 0; y < l.spatial(1); ++y) { - for (cldnn::tensor::value_type x = 0; x < l.spatial(0); ++x) { - auto tensor_coord = cldnn::tensor{{b, f, x, y, z}, 0}; - auto buffer_idx = l.get_linear_offset(tensor_coord); - mem_ptr[buffer_idx] = data[data_idx++]; - } - } - } - } - } - }; - - auto& engine = get_test_engine(); - ov::Shape input0_shape; - ov::Shape input1_shape; - std::vector input1_target_shape; - std::vector input1_output_pattern; - std::vector input0_order; - std::vector input1_order; - ov::Shape beam_table_shape; - cldnn::layout input0_layout; - cldnn::layout input1_layout; - - input0_shape = { BATCH_SIZE, 32, M_SIZE, K_SIZE }; - input1_shape = { N_SIZE, BATCH_SIZE, 2, 1, K_SIZE }; - input1_target_shape = { 1, 1, 1, 16, 1 }; - input1_output_pattern = { 0, 0, 32, K_SIZE }; - input0_order = { 0, 1, 2, 3 }; - input1_order = { 1, 2, 3, 0 }; - - input0_layout = layout{ov::PartialShape::dynamic(input0_shape.size()), data_types::f32, format::bfyx}; - input1_layout = layout{ov::PartialShape::dynamic(input1_shape.size()), data_types::f32, format::bfzyx}; - - auto input0_mem = engine.allocate_memory(layout{ov::PartialShape(input0_shape), data_types::f32, format::bfyx}); - auto input1_mem = engine.allocate_memory(layout{ov::PartialShape(input1_shape), data_types::f32, format::bfzyx}); - - auto input_0_data = rg.generate_random_1d(ov::shape_size(input0_shape), -2, 2); - auto input_1_data = rg.generate_random_1d(ov::shape_size(input1_shape), -2, 2); - - fill_mem(input0_mem, input_0_data); - fill_mem(input1_mem, input_1_data); - - topology topology; - topology.add(input_layout("input0", input0_layout), - input_layout("input1", input1_layout), - gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f32, {}, input1_target_shape, {}, input1_output_pattern, input0_order, input1_order) + gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f32, input0_order, input1_order) ); ExecutionConfig config = get_test_default_config(engine); @@ -1059,12 +922,14 @@ public: cldnn::mem_lock output_ptr(output_mem, get_test_stream()); ov::Shape ref_input0_shape; + ov::Shape ref_input1_unsqueezed_shape; ov::Shape ref_input1_broadcasted_shape; ov::Shape ref_input1_reshaped_shape; ov::Shape ref_input1_shape; ov::Shape ref_output_shape; ref_input0_shape = { BATCH_SIZE, 32, M_SIZE, K_SIZE }; + ref_input1_unsqueezed_shape = { N_SIZE, BATCH_SIZE, 2, 1, K_SIZE }; ref_input1_broadcasted_shape = { N_SIZE, BATCH_SIZE, 2, 16, K_SIZE }; ref_input1_reshaped_shape = { N_SIZE, BATCH_SIZE, 32, K_SIZE }; ref_input1_shape = { BATCH_SIZE, 32, K_SIZE, N_SIZE }; @@ -1087,7 +952,7 @@ public: ov::reference::broadcast(reinterpret_cast(input_1_data.data()), reinterpret_cast(ref_input_1_broadcasted_data.data()), - input1_shape, + ref_input1_unsqueezed_shape, ref_input1_broadcasted_shape, ov::AxisSet({}), sizeof(float)); @@ -1206,7 +1071,7 @@ public: topology topology; topology.add(input_layout("input0", input0_layout), input_layout("input1", input1_layout), - gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f16, {}, {}, {}, {}, input0_order, input1_order, output_order) + gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f16, input0_order, input1_order, output_order) ); ExecutionConfig config = get_test_default_config(engine); @@ -1379,7 +1244,7 @@ public: topology topology; topology.add(input_layout("input0", input0_layout), input_layout("input1", input1_layout), - gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f32, {}, {}, {}, {}, input0_order, input1_order, output_order) + gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f32, input0_order, input1_order, output_order) ); ExecutionConfig config = get_test_default_config(engine); @@ -1493,7 +1358,7 @@ public: topology topology; topology.add(input_layout("input0", input0_layout), input_layout("input1", input1_layout), - gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f16, {}, {}, {}, {}, input0_order, input1_order) + gemm("gemm", { input_info("input0"), input_info("input1") }, data_types::f16, input0_order, input1_order) ); ExecutionConfig config = get_test_default_config(engine); @@ -1668,12 +1533,8 @@ TEST_F(gemm_gpu_tests, transpose_matmul_in1_indirect) { this->test_transpose_indirect(false, false, true); } -TEST_F(gemm_gpu_tests, broadcast_transpose_matmul) { - this->test_broadcast_transpose_matmul(false); -} - -TEST_F(gemm_gpu_tests, broadcast_reshape_transpose_matmul) { - this->test_broadcast_reshape_transpose_matmul(false); +TEST_F(gemm_gpu_tests, unsqueeze_broadcast_reshape_transpose_matmul) { + this->test_unsqueeze_broadcast_reshape_transpose_matmul(false); } TEST_F(gemm_gpu_tests, transpose_matmul_transpose_dynamic_1d) { diff --git a/src/plugins/intel_gpu/tests/unit/transformations/broadcast_reshape_matmul_fusion_test.cpp b/src/plugins/intel_gpu/tests/unit/transformations/unsqueeze_broadcast_reshape_matmul_fusion_test.cpp similarity index 66% rename from src/plugins/intel_gpu/tests/unit/transformations/broadcast_reshape_matmul_fusion_test.cpp rename to src/plugins/intel_gpu/tests/unit/transformations/unsqueeze_broadcast_reshape_matmul_fusion_test.cpp index 8c25f415967..31af9e6f45b 100644 --- a/src/plugins/intel_gpu/tests/unit/transformations/broadcast_reshape_matmul_fusion_test.cpp +++ b/src/plugins/intel_gpu/tests/unit/transformations/unsqueeze_broadcast_reshape_matmul_fusion_test.cpp @@ -9,9 +9,10 @@ #include "openvino/op/broadcast.hpp" #include "openvino/op/constant.hpp" #include "openvino/op/reshape.hpp" +#include "openvino/op/unsqueeze.hpp" #include "intel_gpu/op/gemm.hpp" -#include "plugin/transformations/broadcast_reshape_matmul_fusion.hpp" +#include "plugin/transformations/unsqueeze_broadcast_reshape_matmul_fusion.hpp" #include @@ -22,51 +23,62 @@ namespace ov { namespace test { namespace intel_gpu { -TEST_F(TransformationTestsF, BroadReshapeMatmulFusion1) { +TEST_F(TransformationTestsF, UnsqueezeBroadReshapeMatmulFusion1) { std::vector order_a = {0, 1, 2, 3}; std::vector order_b = {1, 2, 3, 0}; std::vector order_c = {0, 1, 2, 3}; + std::vector axes_b = {-2}; std::vector target_shape_b = {1, 1, 1, 16, 1}; std::vector pattern_b = {0, 0, 32, 32}; { auto input_a = std::make_shared(ov::element::f32, ov::PartialShape::dynamic(4)); - auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 2, 1, 32}); + auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 2, 32}); + auto unsqueeze_b_const = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{1}, axes_b); + auto unsqueeze_b = std::make_shared(input_b, unsqueeze_b_const); auto broadcast_b_const = ov::op::v0::Constant::create(ov::element::i32, ov::Shape{5}, target_shape_b); - auto broadcast_b = std::make_shared(input_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL); + auto broadcast_b = std::make_shared(unsqueeze_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL); auto reshape_b_const = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{4}, pattern_b); auto reshape_b = std::make_shared(broadcast_b, reshape_b_const, true); auto gemm = std::make_shared(input_a, reshape_b, order_a, order_b, order_c, ov::element::undefined); model = std::make_shared(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b }); - manager.register_pass(); + manager.register_pass(); } { auto input_a = std::make_shared(ov::element::f32, ov::PartialShape::dynamic(4)); - auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 2, 1, 32}); - auto gemm = std::make_shared(input_a, input_b, std::vector{}, target_shape_b, std::vector{}, pattern_b, order_a, order_b, order_c, ov::element::undefined); + auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 2, 32}); + auto gemm = std::make_shared(input_a, + input_b, + order_a, + order_b, + order_c, + ov::element::undefined); model_ref = std::make_shared(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b }); comparator.enable(FunctionsComparator::ATTRIBUTES); } } -TEST_F(TransformationTestsF, BroadReshapeMatmulFusion2) { +TEST_F(TransformationTestsF, UnsqueezeBroadReshapeMatmulFusion2) { std::vector order_a = {0, 1, 2, 3}; std::vector order_b = {1, 2, 3, 0}; std::vector order_c = {0, 1, 2, 3}; + std::vector axes_b = {-2}; std::vector target_shape_b = {1, 1, 1, 16, 1}; std::vector pattern_b = {0, 0, -1, 32}; { auto input_a = std::make_shared(ov::element::f32, ov::PartialShape::dynamic(4)); - auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, -1, 1, 32}); + auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, -1, 32}); + auto unsqueeze_b_const = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{1}, axes_b); + auto unsqueeze_b = std::make_shared(input_b, unsqueeze_b_const); auto broadcast_b_const = ov::op::v0::Constant::create(ov::element::i32, ov::Shape{5}, target_shape_b); - auto broadcast_b = std::make_shared(input_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL); + auto broadcast_b = std::make_shared(unsqueeze_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL); auto reshape_b_const = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{4}, pattern_b); auto reshape_b = std::make_shared(broadcast_b, reshape_b_const, true); auto gemm = std::make_shared(input_a, reshape_b, order_a, order_b, order_c, ov::element::undefined); model = std::make_shared(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b }); - manager.register_pass(); + manager.register_pass(); } { model_ref = model->clone(); @@ -74,23 +86,30 @@ TEST_F(TransformationTestsF, BroadReshapeMatmulFusion2) { } } -TEST_F(TransformationTestsF, BroadReshapeMatmulFusion3) { +TEST_F(TransformationTestsF, UnsqueezeBroadReshapeMatmulFusion3) { std::vector order_a = {0, 1, 2, 3}; std::vector order_b = {0, 1, 2, 3}; std::vector order_c = {0, 1, 2, 3}; + std::vector axes_b = {-3}; std::vector target_shape_b = {1, 1, 16, 1, 1}; std::vector pattern_b = {0, 32, 32, 0}; { auto input_a = std::make_shared(ov::element::f32, ov::PartialShape::dynamic(4)); - auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, 2, 1, 32, -1}); + auto input_b = std::make_shared(ov::element::f32, ov::PartialShape{-1, 2, 32, -1}); + auto unsqueeze_b_const = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{1}, axes_b); + auto unsqueeze_b = std::make_shared(input_b, unsqueeze_b_const); auto broadcast_b_const = ov::op::v0::Constant::create(ov::element::i32, ov::Shape{5}, target_shape_b); - auto broadcast_b = std::make_shared(input_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL); + auto broadcast_b = std::make_shared(unsqueeze_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL); auto reshape_b_const = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{4}, pattern_b); auto reshape_b = std::make_shared(broadcast_b, reshape_b_const, true); auto gemm = std::make_shared(input_a, reshape_b, order_a, order_b, order_c, ov::element::undefined); model = std::make_shared(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b }); - manager.register_pass(); + manager.register_pass(); + } + { + model_ref = model->clone(); + comparator.enable(FunctionsComparator::ATTRIBUTES); } }