From e2db808495c4f48af97aeff820fd1e2dcb2105a1 Mon Sep 17 00:00:00 2001 From: Paul Youngsoo Ahn Date: Tue, 15 Aug 2023 05:49:21 +0900 Subject: [PATCH] [GPU] Fix multiple output issue in get_output_layout(#19186) (#19186) --- .../include/intel_gpu/graph/program.hpp | 3 +- src/plugins/intel_gpu/src/graph/program.cpp | 5 +- .../intel_gpu/src/graph/program_node.cpp | 4 +- .../unit/test_cases/arg_max_gpu_test.cpp | 46 ++++++++++++++++++- 4 files changed, 51 insertions(+), 7 deletions(-) diff --git a/src/plugins/intel_gpu/include/intel_gpu/graph/program.hpp b/src/plugins/intel_gpu/include/intel_gpu/graph/program.hpp index 9f52d033a68..dab52844e7e 100644 --- a/src/plugins/intel_gpu/include/intel_gpu/graph/program.hpp +++ b/src/plugins/intel_gpu/include/intel_gpu/graph/program.hpp @@ -139,7 +139,8 @@ public: std::shared_ptr task_executor, bool is_internal); - explicit program(engine& engine); + explicit program(engine& engine, + const ExecutionConfig& config = {}); ~program(); engine& get_engine() const { return _engine; } const ExecutionConfig& get_config() const { return _config; } diff --git a/src/plugins/intel_gpu/src/graph/program.cpp b/src/plugins/intel_gpu/src/graph/program.cpp index 6d69b30fa3d..6f25eb6760c 100644 --- a/src/plugins/intel_gpu/src/graph/program.cpp +++ b/src/plugins/intel_gpu/src/graph/program.cpp @@ -191,10 +191,11 @@ program::program(engine& engine_ref, build_program(is_internal); } -program::program(engine& engine) +program::program(engine& engine, + const ExecutionConfig& config) : _engine(engine), _stream(_engine.create_stream({})), - _config(), + _config(config), processing_order() { init_primitives(); _config.apply_user_properties(_engine.get_device_info()); diff --git a/src/plugins/intel_gpu/src/graph/program_node.cpp b/src/plugins/intel_gpu/src/graph/program_node.cpp index 22ec3e144f5..2fbe366df35 100644 --- a/src/plugins/intel_gpu/src/graph/program_node.cpp +++ b/src/plugins/intel_gpu/src/graph/program_node.cpp @@ -288,8 +288,8 @@ layout program_node::get_output_layout(bool invalidate_users_if_changed, size_t if (valid_output_layouts[idx]) return output_layouts[idx]; - auto new_layout = calc_output_layout(); - set_output_layout(new_layout, invalidate_users_if_changed, idx); + auto new_layouts = calc_output_layouts(); + set_output_layouts(new_layouts, invalidate_users_if_changed); return output_layouts[idx]; } diff --git a/src/plugins/intel_gpu/tests/unit/test_cases/arg_max_gpu_test.cpp b/src/plugins/intel_gpu/tests/unit/test_cases/arg_max_gpu_test.cpp index c665314ab8f..60c3c041611 100644 --- a/src/plugins/intel_gpu/tests/unit/test_cases/arg_max_gpu_test.cpp +++ b/src/plugins/intel_gpu/tests/unit/test_cases/arg_max_gpu_test.cpp @@ -8,11 +8,11 @@ #include #include #include - #include - #include "test_utils.h" +#include "program_wrapper.h" + using namespace cldnn; using namespace ::tests; @@ -930,3 +930,45 @@ TEST(arg_max_min_gpu, dynamic) { ASSERT_FLOAT_EQ(output_ptr[i], i < (out_size / 2) ? 0 : 1); } } + +TEST(arg_max_min_test, check_second_output_data_type) { + auto& engine = get_test_engine(); + + ExecutionConfig config = get_test_default_config(engine); + config.set_property(ov::intel_gpu::allow_new_shape_infer(true)); + + cldnn::program prog(engine, config); + std::vector> input_prims; + std::vector input_prim_ids; + { + auto prim_id = "input"; + static const int32_t x_size = 1, y_size = 1, feature_num = 2000, batch_num = 1; + auto input_static = layout{{batch_num, feature_num, y_size, x_size}, data_types::f16, format::bfyx}; + auto input_layout_prim = std::make_shared(prim_id, input_static); + input_prims.push_back(input_layout_prim); + input_prim_ids.push_back(input_info(prim_id)); + } + { + auto prim_id = "top_k"; + auto top_k_input = layout{{1,1,1,1}, data_types::f16, format::bfyx}; + auto top_k_prim = std::make_shared(prim_id, top_k_input); + input_prims.push_back(top_k_prim); + input_prim_ids.push_back(input_info(prim_id)); + } + + auto arg_max_min_prim = std::make_shared("output", input_prim_ids, + ov::op::TopKMode::MAX, 400, 1, + ov::op::TopKSortType::SORT_VALUES, true, false, padding(), + data_types::f16, 2); + + arg_max_min_prim->output_paddings = {padding(), padding()}; + arg_max_min_prim->output_data_types = {data_types::f16, data_types::i32}; + auto& arg_max_min_node = prog.get_or_create(arg_max_min_prim); + for (auto& prim : input_prims) { + auto& input_layout_node = prog.get_or_create(prim); + program_wrapper::add_connection(prog, input_layout_node, arg_max_min_node); + } + + auto second_output_layout = arg_max_min_node.get_output_layout(false, 1); + ASSERT_EQ(second_output_layout.data_type, data_types::i32); +}