[GPU] Extend gemm to fuse unsqueeze layer (#23734)

### Details:
- Follow up some comments from
https://github.com/openvinotoolkit/openvino/pull/23513
 - Fuse `unsqueeze` layer into `gemm` layer for indirect gemm
    - before : [`kv_cache`] --> [`unsqueeze`] --> `gemm`
    - after : [`kv_cache`] --> `gemm`
 - Simplify fusion pass and logic as `unsqueeze` is fused together

### Tickets:
 - 136567

---------

Signed-off-by: Andrew Park <andrew.park@intel.com>
This commit is contained in:
Andrew Kwangwoong Park 2024-04-17 15:52:32 +09:00 committed by GitHub
parent 6e961cdf89
commit 4b9f92ad83
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
14 changed files with 182 additions and 430 deletions

View File

@ -26,27 +26,12 @@ public:
const std::vector<int64_t>& order_c,
const ov::element::Type output_type = ov::element::undefined);
Gemm(const ov::Output<Node>& A,
const ov::Output<Node>& B,
const std::vector<int32_t>& target_shape_a,
const std::vector<int32_t>& target_shape_b,
const std::vector<int64_t>& output_pattern_a,
const std::vector<int64_t>& output_pattern_b,
const std::vector<int64_t>& order_a,
const std::vector<int64_t>& order_b,
const std::vector<int64_t>& 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<Node> clone_with_new_inputs(const ov::OutputVector& new_args) const override;
std::vector<int32_t> get_input0_broadcast_target_shape() const { return m_target_shape_a; }
std::vector<int32_t> get_input1_broadcast_target_shape() const { return m_target_shape_b; }
std::vector<int64_t> get_input0_reshape_pattern() const { return m_output_pattern_a; }
std::vector<int64_t> get_input1_reshape_pattern() const { return m_output_pattern_b; }
std::vector<int64_t> get_input0_transpose_order() const { return m_order_a; }
std::vector<int64_t> get_input1_transpose_order() const { return m_order_b; }
std::vector<int64_t> get_output_transpose_order() const { return m_order_c; }
@ -59,10 +44,6 @@ public:
}
protected:
std::vector<int32_t> m_target_shape_a;
std::vector<int32_t> m_target_shape_b;
std::vector<int64_t> m_output_pattern_a;
std::vector<int64_t> m_output_pattern_b;
std::vector<int64_t> m_order_a;
std::vector<int64_t> m_order_b;
std::vector<int64_t> m_order_c;
@ -71,10 +52,6 @@ protected:
std::vector<ov::PartialShape> shape_infer(const Gemm* op,
std::vector<ov::PartialShape> input_shapes,
const std::vector<int32_t>& target_shape_a,
const std::vector<int32_t>& target_shape_b,
const std::vector<int64_t>& output_pattern_a,
const std::vector<int64_t>& output_pattern_b,
const std::vector<int64_t>& order_a,
const std::vector<int64_t>& order_b,
const std::vector<int64_t>& order_c);

View File

@ -54,10 +54,6 @@ struct gemm : public primitive_base<gemm> {
: 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> {
gemm(const primitive_id& id,
const std::vector<input_info>& inputs,
const data_types data_type,
const std::vector<int32_t>& input0_broadcast_target_shape = {},
const std::vector<int32_t>& input1_broadcast_target_shape = {},
const std::vector<int64_t>& input0_reshape_pattern = {},
const std::vector<int64_t>& input1_reshape_pattern = {},
const std::vector<int64_t>& input0_transpose_order = {0, 1, 2, 3},
const std::vector<int64_t>& input1_transpose_order = {0, 1, 2, 3},
const std::vector<int64_t>& output_transpose_order = {},
@ -101,10 +93,6 @@ struct gemm : public primitive_base<gemm> {
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<gemm> {
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<gemm> {
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<int32_t> input0_broadcast_target_shape;
/// @brief broadcasted target shape of input 1
std::vector<int32_t> input1_broadcast_target_shape;
/// @brief reshaped output pattern of input 0
std::vector<int64_t> input0_reshape_pattern;
/// @brief reshaped output pattern of input 1
std::vector<int64_t> input1_reshape_pattern;
/// @brief order of input 0
std::vector<int64_t> input0_transpose_order;
/// @brief order of input 1
@ -193,10 +169,6 @@ struct gemm : public primitive_base<gemm> {
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<gemm> {
primitive_base<gemm>::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<gemm> {
primitive_base<gemm>::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;

View File

@ -10,18 +10,6 @@
#include "intel_gpu/op/gemm.hpp"
namespace {
template <typename T, typename DT, typename = typename std::enable_if<std::is_convertible<DT, T>::value>::type>
int find_index_from_vec(const std::vector<T>& vec, const DT value) {
int idx = 0;
for (auto v : vec) {
if (v != static_cast<T>(value))
break;
idx += 1;
}
return idx;
}
} // namespace
namespace cldnn {
GPU_DEFINE_PRIMITIVE_TYPE_ID(gemm)
@ -139,10 +127,6 @@ std::vector<layout> gemm_inst::calc_output_layouts(gemm_node const& node, const
std::vector<ShapeType> 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<layout> gemm_inst::calc_output_layouts<ov::PartialShape>(ge
std::vector<layout> gemm_inst::transform_input_layouts(const std::shared_ptr<const gemm> primitive,
const std::vector<layout>& input_layouts) {
auto get_reshaped_input_shape = [&](const ov::PartialShape& input_pshape,
const std::vector<int32_t>& broadcast_target_shape,
const std::vector<int64_t>& reshape_pattern) {
ov::PartialShape reshaped_input_pshape;
if (broadcast_target_shape.size() > 0 && reshape_pattern.size() > 0) {
std::vector<ov::Dimension> 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<layout> gemm_inst::transform_input_layouts(const std::shared_ptr<con
return transposed_input_pshape;
};
auto reshaped_input0_pshape = get_reshaped_input_shape(input_layouts[0].get_partial_shape(),
primitive->input0_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<layout> 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();

View File

@ -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<int64_t>& 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)) {

View File

@ -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<int32_t>& 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<std::string> 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)),
});
}

View File

@ -19,13 +19,13 @@ struct gemm_params : public base_params {
float beta;
uint32_t transpose_input0;
uint32_t transpose_input1;
std::vector<int32_t> input0_target_shape;
std::vector<int32_t> input1_target_shape;
std::vector<int64_t> input0_output_pattern;
std::vector<int64_t> input1_output_pattern;
std::vector<int64_t> input0_order;
std::vector<int64_t> input1_order;
std::vector<int64_t> 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;

View File

@ -156,23 +156,17 @@ static void CreateGemmOp(ProgramBuilder& p, const std::shared_ptr<ov::op::intern
auto shape_b = op->get_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(),

View File

@ -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<Node>& A,
const std::vector<int64_t>& 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<Node>& A,
const ov::Output<Node>& B,
const std::vector<int32_t>& target_shape_a,
const std::vector<int32_t>& target_shape_b,
const std::vector<int64_t>& output_pattern_a,
const std::vector<int64_t>& output_pattern_b,
const std::vector<int64_t>& order_a,
const std::vector<int64_t>& order_b,
const std::vector<int64_t>& 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<Node>& A,
std::shared_ptr<ov::Node> Gemm::clone_with_new_inputs(const ov::OutputVector& new_args) const {
check_new_args_count(this, new_args);
return std::make_shared<Gemm>(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<Gemm>(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<ov::PartialShape>{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<ov::PartialShape> shape_infer(const Gemm* op,
std::vector<ov::PartialShape> input_shapes,
const std::vector<int32_t>& target_shape_a,
const std::vector<int32_t>& target_shape_b,
const std::vector<int64_t>& output_pattern_a,
const std::vector<int64_t>& output_pattern_b,
const std::vector<int64_t>& order_a,
const std::vector<int64_t>& order_b,
const std::vector<int64_t>& 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<int32_t>& target_shape) {
ov::op::v3::Broadcast broadcast;
auto tshape = target_shape;
broadcast.set_broadcast_spec(ov::op::BroadcastType::BIDIRECTIONAL);
std::unordered_map<size_t, ov::Tensor> const_data;
const_data.emplace(1, ov::Tensor(ov::element::i32, ov::Shape{tshape.size()}, static_cast<void*>(tshape.data())));
return ov::op::v3::shape_infer(&broadcast,
std::vector<ov::PartialShape>{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<int64_t>& output_pattern) {
ov::op::v1::Reshape reshape;
auto opattern = output_pattern;
reshape.set_special_zero(true);
std::unordered_map<size_t, ov::Tensor> const_data;
const_data.emplace(1, ov::Tensor(ov::element::i64, ov::Shape{opattern.size()}, static_cast<void*>(opattern.data())));
return ov::op::v1::shape_infer(&reshape,
std::vector<ov::PartialShape>{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<int64_t>& order) {
auto shape_transposed = ov::PartialShape::dynamic(shape.rank());
// transposed shape
auto transpose_pshape = [](const ov::PartialShape pshape, const std::vector<int64_t>& 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<const ov::op::v0::MatMul*>(op), std::vector<ov::PartialShape>{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] };
}

View File

@ -50,10 +50,6 @@ void IndirectGemm::validate_and_infer_types() {
auto out_shapes = shape_infer(this,
std::vector<ov::PartialShape>{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);

View File

@ -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<ov::Node>& 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<ov::op::v0::Constant>();
auto unsqueeze_a_m = wrap_type<ov::op::v0::Unsqueeze>({input_a_m, unsqueeze_a_axes_m}, consumers_count(1));
auto unsqueeze_b_axes_m = wrap_type<ov::op::v0::Constant>();
auto unsqueeze_b_m = wrap_type<ov::op::v0::Unsqueeze>({input_b_m, unsqueeze_b_axes_m}, consumers_count(1));
auto broadcast_a_target_shape_m = wrap_type<ov::op::v0::Constant>();
auto broadcast_a_m = wrap_type<ov::op::v3::Broadcast>({input_a_m, broadcast_a_target_shape_m}, broadcast_rank_equals_and_has_static_dims);
auto broadcast_a_m = wrap_type<ov::op::v3::Broadcast>({unsqueeze_a_m, broadcast_a_target_shape_m}, broadcast_rank_equals_and_has_static_dims);
auto broadcast_b_target_shape_m = wrap_type<ov::op::v0::Constant>();
auto broadcast_b_m = wrap_type<ov::op::v3::Broadcast>({input_b_m, broadcast_b_target_shape_m}, broadcast_rank_equals_and_has_static_dims);
auto broadcast_b_m = wrap_type<ov::op::v3::Broadcast>({unsqueeze_b_m, broadcast_b_target_shape_m}, broadcast_rank_equals_and_has_static_dims);
auto reshape_a_pattern_m = wrap_type<ov::op::v0::Constant>();
auto reshape_a_m = wrap_type<ov::op::v1::Reshape>({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<int32_t>();
auto target_shape_b = std::vector<int32_t>();
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<ov::op::v0::Unsqueeze>(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<ov::op::v0::Unsqueeze>(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<int32_t>();
auto target_shape_b = std::vector<int32_t>();
auto valid_broadcast_target_shape = [](const std::vector<int32_t>& 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<ov::op::v3::Broadcast>(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<int32_t>();
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<ov::op::v3::Broadcast>(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<int32_t>();
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<int64_t>();
@ -123,10 +141,6 @@ BroadcastReshapeMatmulFusion::BroadcastReshapeMatmulFusion() {
auto gemm = std::make_shared<op::Gemm>(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<ov::pass::pattern::Matcher>(matmul_m, "BroadcastReshapeMatmulFusion");
auto m = std::make_shared<ov::pass::pattern::Matcher>(matmul_m, "UnsqueezeBroadcastReshapeMatmulFusion");
this->register_matcher(m, callback);
}

View File

@ -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

View File

@ -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<ov::Model> func) {
manager.register_pass<ov::intel_gpu::RMSFusion>(device_info.max_work_group_size);
manager.register_pass<ov::intel_gpu::KVCacheFusion>();
manager.register_pass<ov::intel_gpu::FullyConnectedConvertFusion>();
if (!device_info.supports_immad)
if (!device_info.supports_immad) {
manager.register_pass<ov::intel_gpu::TransposeMatMulFusion>();
manager.register_pass<ov::intel_gpu::UnsqueezeBroadcastReshapeMatmulFusion>();
}
manager.register_pass<ov::intel_gpu::SwiGLUFusion>();
manager.register_pass<ov::intel_gpu::IndirectKVCache>();
manager.register_pass<ov::intel_gpu::ConvertConvolutionToInternal>();
if (!device_info.supports_immad)
manager.register_pass<ov::intel_gpu::BroadcastReshapeMatmulFusion>();
const size_t zp_pad_size = device_info.supports_immad ? 16 : 32;
manager.register_pass<ov::intel_gpu::BroadcastAndPadZeroPointBuffers>(zp_pad_size);

View File

@ -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<int32_t> input1_target_shape;
std::vector<int64_t> input0_order;
std::vector<int64_t> 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<float> 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<float> ref_out_data;
ref_out_data.resize(ov::shape_size(ref_output_shape));
std::vector<float> ref_input_0_data(input_0_data.size());
std::vector<float> ref_input_1_broadcasted_data(ov::shape_size(ref_input1_broadcasted_shape));
std::vector<float> 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<const char*>(input_1_data.data()),
reinterpret_cast<char*>(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<float>(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<float>& data) {
cldnn::mem_lock<float> 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<int32_t> input1_target_shape;
std::vector<int64_t> input1_output_pattern;
std::vector<int64_t> input0_order;
std::vector<int64_t> 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<float>(ov::shape_size(input0_shape), -2, 2);
auto input_1_data = rg.generate_random_1d<float>(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<float> 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<const char*>(input_1_data.data()),
reinterpret_cast<char*>(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) {

View File

@ -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 <memory>
@ -22,51 +23,62 @@ namespace ov {
namespace test {
namespace intel_gpu {
TEST_F(TransformationTestsF, BroadReshapeMatmulFusion1) {
TEST_F(TransformationTestsF, UnsqueezeBroadReshapeMatmulFusion1) {
std::vector<int64_t> order_a = {0, 1, 2, 3};
std::vector<int64_t> order_b = {1, 2, 3, 0};
std::vector<int64_t> order_c = {0, 1, 2, 3};
std::vector<int64_t> axes_b = {-2};
std::vector<int32_t> target_shape_b = {1, 1, 1, 16, 1};
std::vector<int64_t> pattern_b = {0, 0, 32, 32};
{
auto input_a = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape::dynamic(4));
auto input_b = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape{-1, -1, 2, 1, 32});
auto input_b = std::make_shared<ov::op::v0::Parameter>(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<ov::op::v0::Unsqueeze>(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<ov::op::v3::Broadcast>(input_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL);
auto broadcast_b = std::make_shared<ov::op::v3::Broadcast>(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<ov::op::v1::Reshape>(broadcast_b, reshape_b_const, true);
auto gemm = std::make_shared<ov::intel_gpu::op::Gemm>(input_a, reshape_b, order_a, order_b, order_c, ov::element::undefined);
model = std::make_shared<ov::Model>(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b });
manager.register_pass<BroadcastReshapeMatmulFusion>();
manager.register_pass<UnsqueezeBroadcastReshapeMatmulFusion>();
}
{
auto input_a = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape::dynamic(4));
auto input_b = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape{-1, -1, 2, 1, 32});
auto gemm = std::make_shared<ov::intel_gpu::op::Gemm>(input_a, input_b, std::vector<int32_t>{}, target_shape_b, std::vector<int64_t>{}, pattern_b, order_a, order_b, order_c, ov::element::undefined);
auto input_b = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape{-1, -1, 2, 32});
auto gemm = std::make_shared<ov::intel_gpu::op::Gemm>(input_a,
input_b,
order_a,
order_b,
order_c,
ov::element::undefined);
model_ref = std::make_shared<ov::Model>(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b });
comparator.enable(FunctionsComparator::ATTRIBUTES);
}
}
TEST_F(TransformationTestsF, BroadReshapeMatmulFusion2) {
TEST_F(TransformationTestsF, UnsqueezeBroadReshapeMatmulFusion2) {
std::vector<int64_t> order_a = {0, 1, 2, 3};
std::vector<int64_t> order_b = {1, 2, 3, 0};
std::vector<int64_t> order_c = {0, 1, 2, 3};
std::vector<int64_t> axes_b = {-2};
std::vector<int32_t> target_shape_b = {1, 1, 1, 16, 1};
std::vector<int64_t> pattern_b = {0, 0, -1, 32};
{
auto input_a = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape::dynamic(4));
auto input_b = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape{-1, -1, -1, 1, 32});
auto input_b = std::make_shared<ov::op::v0::Parameter>(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<ov::op::v0::Unsqueeze>(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<ov::op::v3::Broadcast>(input_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL);
auto broadcast_b = std::make_shared<ov::op::v3::Broadcast>(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<ov::op::v1::Reshape>(broadcast_b, reshape_b_const, true);
auto gemm = std::make_shared<ov::intel_gpu::op::Gemm>(input_a, reshape_b, order_a, order_b, order_c, ov::element::undefined);
model = std::make_shared<ov::Model>(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b });
manager.register_pass<BroadcastReshapeMatmulFusion>();
manager.register_pass<UnsqueezeBroadcastReshapeMatmulFusion>();
}
{
model_ref = model->clone();
@ -74,23 +86,30 @@ TEST_F(TransformationTestsF, BroadReshapeMatmulFusion2) {
}
}
TEST_F(TransformationTestsF, BroadReshapeMatmulFusion3) {
TEST_F(TransformationTestsF, UnsqueezeBroadReshapeMatmulFusion3) {
std::vector<int64_t> order_a = {0, 1, 2, 3};
std::vector<int64_t> order_b = {0, 1, 2, 3};
std::vector<int64_t> order_c = {0, 1, 2, 3};
std::vector<int64_t> axes_b = {-3};
std::vector<int32_t> target_shape_b = {1, 1, 16, 1, 1};
std::vector<int64_t> pattern_b = {0, 32, 32, 0};
{
auto input_a = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape::dynamic(4));
auto input_b = std::make_shared<ov::op::v0::Parameter>(ov::element::f32, ov::PartialShape{-1, 2, 1, 32, -1});
auto input_b = std::make_shared<ov::op::v0::Parameter>(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<ov::op::v0::Unsqueeze>(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<ov::op::v3::Broadcast>(input_b, broadcast_b_const, ov::op::BroadcastType::BIDIRECTIONAL);
auto broadcast_b = std::make_shared<ov::op::v3::Broadcast>(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<ov::op::v1::Reshape>(broadcast_b, reshape_b_const, true);
auto gemm = std::make_shared<ov::intel_gpu::op::Gemm>(input_a, reshape_b, order_a, order_b, order_c, ov::element::undefined);
model = std::make_shared<ov::Model>(ov::NodeVector{ gemm }, ov::ParameterVector{ input_a, input_b });
manager.register_pass<BroadcastReshapeMatmulFusion>();
manager.register_pass<UnsqueezeBroadcastReshapeMatmulFusion>();
}
{
model_ref = model->clone();
comparator.enable(FunctionsComparator::ATTRIBUTES);
}
}