From 902c370ca17bd2333e7d2f82ae1274b66762e6a7 Mon Sep 17 00:00:00 2001 From: Vladimir Paramuzov Date: Wed, 3 Apr 2024 03:14:02 -0700 Subject: [PATCH] [GPU] Remove paddings on set_state() call (#23828) ### Details: - Fix accuracy issue due to wrong pad size after set_state call --- .../intel_gpu/src/plugin/variable_state.cpp | 2 ++ .../subgraph_tests/dynamic/kv_cache.cpp | 15 ++++++++++++++- 2 files changed, 16 insertions(+), 1 deletion(-) diff --git a/src/plugins/intel_gpu/src/plugin/variable_state.cpp b/src/plugins/intel_gpu/src/plugin/variable_state.cpp index 2c85fcf4663..19c8c20016b 100644 --- a/src/plugins/intel_gpu/src/plugin/variable_state.cpp +++ b/src/plugins/intel_gpu/src/plugin/variable_state.cpp @@ -58,6 +58,8 @@ void VariableState::set_layout(const cldnn::layout& new_layout) { void VariableState::set_state(const ov::SoPtr& state) { m_layout.set_partial_shape(state->get_shape()); + size_t rank = state->get_shape().size(); + m_layout.data_padding = cldnn::padding(std::vector(rank, 0), std::vector(rank, 0), 0, m_layout.data_padding.get_dynamic_pad_dims()); update_device_buffer(); convert_and_copy(state._ptr.get(), m_memory, m_context->get_engine().get_service_stream()); set(); diff --git a/src/plugins/intel_gpu/tests/functional/subgraph_tests/dynamic/kv_cache.cpp b/src/plugins/intel_gpu/tests/functional/subgraph_tests/dynamic/kv_cache.cpp index f9fb504ea12..e5461ca96d7 100644 --- a/src/plugins/intel_gpu/tests/functional/subgraph_tests/dynamic/kv_cache.cpp +++ b/src/plugins/intel_gpu/tests/functional/subgraph_tests/dynamic/kv_cache.cpp @@ -258,7 +258,8 @@ class KVCacheTests: public ::testing::Test { int64_t concat_axis = 2, ov::element::Type model_element_type = ov::element::f16, size_t num_iter = 10, - size_t num_groups = 1) { + size_t num_groups = 1, + bool set_state_on_each_iter = false) { #if defined(ANDROID) GTEST_SKIP(); #endif @@ -437,6 +438,14 @@ class KVCacheTests: public ::testing::Test { infer_request.infer(); compare_tensors({ ref_results[1] }, {matmul_out}); + + if (set_state_on_each_iter) { + auto state = infer_request.query_state()[0].get_state(); + compare_tensors({ ref_kv_cache }, {state}); + infer_request.query_state()[0].set_state(state); + auto state_1 = infer_request.query_state()[0].get_state(); + compare_tensors({ ref_kv_cache }, {state_1}); + } } auto state = infer_request.query_state()[0].get_state(); @@ -494,4 +503,8 @@ TEST_F(KVCacheTests, smoke_multipleIterations_stateful_same_shape_after_reset) { this->test_smoke_multipleIterations_stateful(false, false, false, 1, 2, ov::element::f16, 0); } +TEST_F(KVCacheTests, smoke_multipleIterations_stateful_with_set_state) { + this->test_smoke_multipleIterations_stateful(false, true, true, 1, 2, ov::element::f16, 5, 1, true); +} + } // namespace