From 18a9c772b426c9f422c52a703aabdb6c21f1a870 Mon Sep 17 00:00:00 2001 From: Maksim Kutakov Date: Mon, 18 Dec 2023 20:16:32 +0100 Subject: [PATCH] [CPU] Plugin optimizations for LLM beam search via states (#21642) --- src/plugins/intel_cpu/src/edge.cpp | 1 + src/plugins/intel_cpu/src/graph.cpp | 8 +- src/plugins/intel_cpu/src/graph_optimizer.cpp | 40 +- .../memory_desc/cpu_blocked_memory_desc.cpp | 3 - src/plugins/intel_cpu/src/memory_state.cpp | 165 +++-- src/plugins/intel_cpu/src/memory_state.h | 52 +- .../nodes/kernels/scaled_attn/attn_memcpy.cpp | 34 +- src/plugins/intel_cpu/src/nodes/memory.cpp | 598 +++++++++++------- src/plugins/intel_cpu/src/nodes/memory.hpp | 115 +++- .../intel_cpu/src/nodes/scaled_attn.cpp | 540 ++++++++++++---- src/plugins/intel_cpu/src/nodes/scaled_attn.h | 31 +- .../cpu_opset/common/pass/rope_fusion.cpp | 4 +- .../common/pass/stateful_sdpa_fusion.cpp | 156 ++++- .../pass/stateful_transpose_sdpa_fusion.cpp | 176 ------ .../pass/stateful_transpose_sdpa_fusion.hpp | 18 - .../transformation_pipeline.cpp | 2 - .../intel_cpu/src/utils/plain_tensor.hpp | 11 +- .../src/concat_multiple_query_sdp.cpp | 84 ++- .../subgraph_tests/src/concat_sdp.cpp | 36 +- .../src/concat_transpose_sdp_transpose.cpp | 220 +++++-- .../tests/unit/graph/scaled_attn.cpp | 177 ------ .../transformations/state_concat_sdpa.cpp | 7 +- 22 files changed, 1543 insertions(+), 935 deletions(-) delete mode 100644 src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.cpp delete mode 100644 src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.hpp delete mode 100644 src/plugins/intel_cpu/tests/unit/graph/scaled_attn.cpp diff --git a/src/plugins/intel_cpu/src/edge.cpp b/src/plugins/intel_cpu/src/edge.cpp index 010806e3122..a090af2273e 100644 --- a/src/plugins/intel_cpu/src/edge.cpp +++ b/src/plugins/intel_cpu/src/edge.cpp @@ -458,6 +458,7 @@ void Edge::init() { changeStatus(Status::NeedAllocation); } else { if (Type::Input == edgePtr->getParent()->getType() && + Type::MemoryInput != getParent()->getType() && edgePtr->getParent()->isConstant() && !edgePtr->getChild()->isConstant()) { changeStatus(Status::NeedAllocation); diff --git a/src/plugins/intel_cpu/src/graph.cpp b/src/plugins/intel_cpu/src/graph.cpp index d8c0d9ffc55..0f8bdfea859 100644 --- a/src/plugins/intel_cpu/src/graph.cpp +++ b/src/plugins/intel_cpu/src/graph.cpp @@ -875,9 +875,13 @@ void Graph::Allocate() { //resolve inplace dead end nodes for (const auto& edge : graphEdges) { if (edge->getStatus() == Edge::Status::Uninitialized) { - if (one_of(edge->getParent()->getType(), Type::Input, Type::MemoryInput) && edge->inPlace(Edge::LOOK_UP)) { + if (edge->getParent()->getParentEdges().empty() && + one_of(edge->getParent()->getType(), Type::Input, Type::MemoryInput) && + edge->inPlace(Edge::LOOK_UP)) { edge->getParent()->resolveInPlaceEdges(Edge::LOOK_UP); - } else if (one_of(edge->getChild()->getType(), Type::Output, Type::MemoryOutput) && edge->inPlace(Edge::LOOK_DOWN)) { + } else if (edge->getChild()->getChildEdges().empty() && + one_of(edge->getChild()->getType(), Type::Output, Type::MemoryOutput) && + edge->inPlace(Edge::LOOK_DOWN)) { edge->getChild()->resolveInPlaceEdges(Edge::LOOK_DOWN); } } diff --git a/src/plugins/intel_cpu/src/graph_optimizer.cpp b/src/plugins/intel_cpu/src/graph_optimizer.cpp index c79f0ed2aa6..63c7e1cc8d1 100644 --- a/src/plugins/intel_cpu/src/graph_optimizer.cpp +++ b/src/plugins/intel_cpu/src/graph_optimizer.cpp @@ -23,6 +23,7 @@ #include "nodes/input.h" #include "nodes/rnn.h" #include "nodes/memory.hpp" +#include "nodes/scaled_attn.h" #include "nodes/common/cpu_convert.h" #include "onednn/dnnl.h" @@ -2837,6 +2838,20 @@ void GraphOptimizer::MatchSdpaKvCache(Graph &graph) { input_prc = ov::optional(node->getOriginalInputPrecisionAtPort(0)); } + //search for SDPA + std::shared_ptr sdpa; + for (auto&& edge : node->getChildEdgesAtPort(0)) { + auto child = edge->getChild(); + if (Type::ScaledDotProductAttention == child->getType()) { + sdpa = std::dynamic_pointer_cast(child); + if (sdpa) { + break; + } else { + OPENVINO_THROW("Couldn't cast node", child->getName(), " to ScaledDotProductAttention type"); + } + } + } + auto memInputSdpa = std::make_shared( memInputNode->getId(), memInputNode->getName(), @@ -2845,7 +2860,8 @@ void GraphOptimizer::MatchSdpaKvCache(Graph &graph) { memInputNode->getOriginalOutputPrecisionAtPort(0), graph.getGraphContext(), input_shape, - input_prc); + input_prc, + sdpa); if (!memInputNode->getParentEdges().empty()) { auto parentEdge = memInputNode->getParentEdgeAt(0); @@ -2862,12 +2878,28 @@ void GraphOptimizer::MatchSdpaKvCache(Graph &graph) { graph.RemoveEdge(edge); } - //link with memory output + //create a stub memory output auto& memOutput = memInputNode->getOutputNode(); - memInputSdpa->registerOutputNode(&memOutput); + + auto memOutputStub = std::make_shared( + memOutput.getId(), + memOutput.getName(), + memOutput.getTypeStr(), + memOutput.getInputShapeAtPort(0), + memOutput.getOriginalInputPrecisionAtPort(0), + graph.getGraphContext()); + + auto memOutputEdge = memOutput.getParentEdgeAt(0); + auto newEdge = + std::make_shared(sdpa, memOutputStub, memOutputEdge->getInputNum(), 0); + memOutputStub->addEdge(newEdge); + graph.GetEdges().push_back(newEdge); + graph.RemoveEdge(memOutputEdge); + + memInputSdpa->registerOutputNode(memOutputStub.get()); graph.GetNodes().push_back(memInputSdpa); - graph.DropNode(memInputNode); + graph.GetNodes().push_back(memOutputStub); } } diff --git a/src/plugins/intel_cpu/src/memory_desc/cpu_blocked_memory_desc.cpp b/src/plugins/intel_cpu/src/memory_desc/cpu_blocked_memory_desc.cpp index c119628c410..ab3df84a06f 100644 --- a/src/plugins/intel_cpu/src/memory_desc/cpu_blocked_memory_desc.cpp +++ b/src/plugins/intel_cpu/src/memory_desc/cpu_blocked_memory_desc.cpp @@ -60,9 +60,6 @@ CpuBlockedMemoryDesc::CpuBlockedMemoryDesc(ov::element::Type prc, const Shape& s } } } else { - if (shape.hasZeroDims() && std::any_of(strides.begin(), strides.end(), [](size_t stride) { return stride != 0; } )) { - OPENVINO_THROW("Can't create CpuBlockedMemoryDesc with zero dim, but with non zero strides"); - } this->strides = strides; } diff --git a/src/plugins/intel_cpu/src/memory_state.cpp b/src/plugins/intel_cpu/src/memory_state.cpp index 968d95222cb..371e60c0f0e 100644 --- a/src/plugins/intel_cpu/src/memory_state.cpp +++ b/src/plugins/intel_cpu/src/memory_state.cpp @@ -8,6 +8,9 @@ #include "dnnl_extension_utils.h" #include "blob_factory.hpp" #include "cpu_tensor.h" +#include "utils/plain_tensor.hpp" +#include "openvino/core/parallel.hpp" +#include "nodes/common/cpu_convert.h" using namespace InferenceEngine; @@ -34,9 +37,8 @@ const dnnl::engine& VariableStateBase::get_engine() { return eng; } -void VariableStateBase::set_state(const ov::SoPtr& state) { - m_state = state; // simply to extend the lifetime - auto state_desc = MemoryDescUtils::generateCpuBlockedMemoryDesc(m_state); +void VariableStateBase::set_state_impl(const ov::SoPtr& state) { + auto state_desc = MemoryDescUtils::generateCpuBlockedMemoryDesc(state); const auto& shape = state_desc->getShape(); @@ -45,13 +47,18 @@ void VariableStateBase::set_state(const ov::SoPtr& state) { input_mem()->redefineDesc(new_desc); } - auto src = m_state->data(); + auto src = state->data(); Memory mem(get_engine(), state_desc, src); input_mem()->load(mem); reset_state_flag = false; } +void VariableStateBase::set_state(const ov::SoPtr& state) { + set_state_impl(state); + reset_state_flag = false; +} + ov::SoPtr VariableStateBase::get_state() const { const auto& current_dims = internal_state_mem()->getStaticDims(); auto current_ext_desc = m_external_desc->cloneWithNewDims(current_dims); @@ -146,50 +153,124 @@ MemoryPtr VariableStateDoubleBuffer::internal_state_mem() const { return prime_mem(); } -VariableStateSingleBuffer::VariableStateSingleBuffer(const std::string& name, - const MemoryPtr& buffer, - const MemoryDescPtr& external_desc) : - VariableStateBase(name, external_desc) { - OPENVINO_ASSERT(buffer); - m_internal_mem = buffer; - m_internal_desc = m_internal_mem->getDescPtr(); - auto&& shape = m_internal_desc->getShape(); - //TODO what if by some reason we already have internal static state while the node is dynamic, is it even possible? +VariableStateKVcache::VariableStateKVcache( + const std::string& name, + const MemoryDescPtr& external_desc, + const BlockedMemoryDescPtr& dense_internal_desc) : + VariableStateBase(name, external_desc), m_dense_internal_desc(dense_internal_desc) { + auto&& shape = external_desc->getShape(); - if (shape.isStatic()) { - m_internal_mem->nullify(); - } else { - //in the case of the original desc has dynamic shape we create an empty tensor - auto new_desc = to_static(m_internal_desc); - m_internal_mem->redefineDesc(new_desc); + OPENVINO_ASSERT(shape.isDynamic(), "VariableStateKVcache is unexpectedly initalized with a static tensor"); +} + +ov::SoPtr VariableStateKVcache::get_state() const { + OPENVINO_ASSERT(m_internal_mem && m_hidden_state, "KVState internal memory is not initialized"); + OPENVINO_ASSERT(!is_reset_state(), "KVState is undefined after reset"); + auto actual_internal_desc = m_internal_mem->getDescWithType(); + auto&& dims = actual_internal_desc->getShape().getStaticDims(); + + auto actual_external_desc = get_external_desc()->cloneWithNewDims(dims); + auto external_mem = std::make_shared(get_engine(), actual_external_desc); + + // let's assume 4th rank KV tensors. This may be extended later + OPENVINO_ASSERT(actual_internal_desc->getShape().getRank() == 4); + OPENVINO_ASSERT(actual_external_desc->getShape().getRank() == 4); + + auto&& actual_internal_order = actual_internal_desc->getOrder(); + //sanity check + OPENVINO_ASSERT(actual_internal_order == m_dense_internal_desc->getOrder()); + + PlainTensor output, pastkv, beam_table; + output.reset(external_mem); + beam_table.reset(m_hidden_state); + pastkv.reset(m_internal_mem); + output = output.permute(actual_internal_order); + pastkv = pastkv.permute(actual_internal_order); + // S should be always the last dimension + OPENVINO_ASSERT(pastkv.stride(3) == 1 && output.stride(3) == 1); + auto B = pastkv.size(0); + auto H = pastkv.size(1); + auto L0 = pastkv.size(2); + auto S = pastkv.size(3); + parallel_for3d(B, H, L0, [&](size_t b, size_t h, size_t m) { + auto b_kv = static_cast(beam_table.at({b, m})); + cpu_convert(&pastkv.at({b_kv, h, m}), + &output.at({b, h, m}), + pastkv.m_dt, + output.m_dt, + S); + }); + + return std::make_shared(external_mem); +} + +void VariableStateKVcache::set_state_impl(const ov::SoPtr& state) { + //1. reset the memory object + m_state = state; // simply to extend the lifetime + auto state_desc = MemoryDescUtils::generateCpuBlockedMemoryDesc(m_state); + + //May be optimized by reusing the state tensor underlining memory pointer, but corner cases should be considered + auto dense_internal_desc = m_dense_internal_desc->cloneWithNewDims(state_desc->getShape().getStaticDims()); + + m_internal_mem = std::make_shared(get_engine(), dense_internal_desc); + Memory external_mem(get_engine(), state_desc, m_state->data()); + + m_internal_mem->load(external_mem); + + //2. Reset the beam search table + auto&& state_dims = dense_internal_desc->getShape().getStaticDims(); + auto&& order = m_dense_internal_desc->getOrder(); + + const size_t size_B = state_dims[order.at(0)]; + const size_t size_L = state_dims[order.at(2)]; + auto mem_desc = + std::make_shared(ov::element::i32, Shape{size_B, size_L}); + + m_hidden_state = std::make_shared(get_engine(), mem_desc); + auto buff = reinterpret_cast(m_hidden_state->getData()); + for (size_t i = 0; i < size_B; ++i) { + for (size_t j = 0; j < size_L; ++j) { + buff[i * size_L + j] = i; + } } + m_internal_mem_max_size = dense_internal_desc->getCurrentMemSize() / dense_internal_desc->getPrecision().size(); + m_hidden_state_max_size = mem_desc->getCurrentMemSize() / mem_desc->getPrecision().size(); } -void VariableStateSingleBuffer::reset_impl() { - auto new_desc = to_static(m_internal_desc); - m_internal_mem->redefineDesc(new_desc); - m_internal_mem->nullify(); -} - -MemoryPtr VariableStateSingleBuffer::input_mem() { - return m_internal_mem; -} - -MemoryPtr VariableStateSingleBuffer::output_mem() { - return m_internal_mem; -} - -MemoryDescPtr VariableStateSingleBuffer::internal_desc() const { - return m_internal_desc; -} - -MemoryPtr VariableStateSingleBuffer::internal_state_mem() const { - return m_internal_mem; -} - -void VariableStateSingleBuffer::commit_impl() { +void VariableStateKVcache::reset_impl() { //nothing to do } +void VariableStateKVcache::commit_impl() { + //nothing to do +} + +MemoryPtr VariableStateKVcache::input_mem() { + return m_internal_mem; +} + +MemoryPtr VariableStateKVcache::output_mem() { + return m_internal_mem; +} + +MemoryDescPtr VariableStateKVcache::internal_desc() const { + return m_internal_mem->getDescPtr(); //since we don't store initial one +} + +MemoryPtr VariableStateKVcache::internal_state_mem() const { + return m_internal_mem; +} + +void VariableStateKVcache::assign_internal_state(const MemoryPtr& mem) { + m_internal_mem = mem; +} + +MemoryPtr VariableStateKVcache::hidden_state_mem() const { + return m_hidden_state; +} + +void VariableStateKVcache::assign_hidden_state(const MemoryPtr& mem) { + m_hidden_state = mem; +} } // namespace intel_cpu } // namespace ov diff --git a/src/plugins/intel_cpu/src/memory_state.h b/src/plugins/intel_cpu/src/memory_state.h index 29ed964f647..ef407bddaa8 100644 --- a/src/plugins/intel_cpu/src/memory_state.h +++ b/src/plugins/intel_cpu/src/memory_state.h @@ -34,7 +34,7 @@ public: VariableStateBase(const std::string& name, const MemoryDescPtr& external_desc); //ov::IVariableState - void set_state(const ov::SoPtr& state) override; + void set_state(const ov::SoPtr& state) override final; // NOLINT ov::SoPtr get_state() const override; void reset() override final; // NOLINT bool is_reset_state() const override final; // NOLINT @@ -44,10 +44,15 @@ protected: virtual MemoryPtr internal_state_mem() const = 0; virtual void reset_impl() = 0; virtual void commit_impl() = 0; + virtual void set_state_impl(const ov::SoPtr& state); static MemoryDescPtr to_static(const MemoryDescPtr& desc); static const dnnl::engine& get_engine(); + MemoryDescPtr get_external_desc() const { + return m_external_desc; + } + private: MemoryDescPtr m_external_desc; bool reset_state_flag = true; @@ -93,26 +98,55 @@ private: size_t buffer_num = 0; }; -class VariableStateSingleBuffer : public VariableStateBase { +class VariableStateKVcache : public VariableStateBase { public: - VariableStateSingleBuffer(const std::string& name, - const MemoryPtr& buffer, - const MemoryDescPtr& external_desc); + VariableStateKVcache(const std::string& name, + const MemoryDescPtr& external_desc, + const BlockedMemoryDescPtr& dense_internal_desc); + //ov::IVariableState + ov::SoPtr get_state() const override; + + //ov::intel_cpu::VariableStateBase MemoryPtr input_mem() override; MemoryPtr output_mem() override; MemoryDescPtr internal_desc() const override; + MemoryPtr internal_state_mem() const override; + void assign_internal_state(const MemoryPtr& mem); + + MemoryPtr hidden_state_mem() const; + void assign_hidden_state(const MemoryPtr& mem); + + // size in elements count + size_t internal_state_max_size() const { + return m_internal_mem_max_size; + } + void assign_internal_state_max_size(size_t max_size) { + m_internal_mem_max_size = max_size; + } + + size_t hidden_state_max_size() const { + return m_hidden_state_max_size; + } + void assign_hidden_state_max_size(size_t max_size) { + m_hidden_state_max_size = max_size; + } + private: //ov::intel_cpu::VariableStateBase + void set_state_impl(const ov::SoPtr& state) override; void reset_impl() override; void commit_impl() override; - MemoryPtr internal_state_mem() const override; - private: - MemoryDescPtr m_internal_desc; //mem desc required by the graph internal tensor - MemoryPtr m_internal_mem; + MemoryPtr m_internal_mem; // kv cache + MemoryPtr m_hidden_state; // beam access table + size_t m_internal_mem_max_size = 0; + size_t m_hidden_state_max_size = 0; + + // this desc stores the internal prc and axis permutation + BlockedMemoryDescPtr m_dense_internal_desc; }; using MemStatePtr = std::shared_ptr; diff --git a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/attn_memcpy.cpp b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/attn_memcpy.cpp index 6c2b35b94f4..a2d653c69fc 100644 --- a/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/attn_memcpy.cpp +++ b/src/plugins/intel_cpu/src/nodes/kernels/scaled_attn/attn_memcpy.cpp @@ -61,19 +61,18 @@ void attn_memcpy_kernel(const ov::intel_cpu::PlainTensor& k_input, }); } -template -void attn_memcpy_kernel(const ov::intel_cpu::PlainTensor& k_input, - const ov::intel_cpu::PlainTensor& v_input, - const ov::intel_cpu::PlainTensor& past_k_output, - const ov::intel_cpu::PlainTensor& past_v_output) { +static void attn_memcpy_kernel(const ov::intel_cpu::PlainTensor& k_input, + const ov::intel_cpu::PlainTensor& v_input, + const ov::intel_cpu::PlainTensor& past_k_output, + const ov::intel_cpu::PlainTensor& past_v_output) { size_t B = k_input.m_dims[0], H = k_input.m_dims[1], L1 = k_input.m_dims[2], S = k_input.m_dims[3]; parallel_for3d(B, H, L1, [&](size_t b, size_t h, size_t m) { - memcpy(&past_k_output.at({b, h, m, 0}), - &k_input.at({b, h, m, 0}), - S * sizeof(T)); - memcpy(&past_v_output.at({b, h, m, 0}), - &v_input.at({b, h, m, 0}), - S * sizeof(T)); + std::memcpy(&past_k_output.at({b, h, m, 0}), + &k_input.at({b, h, m, 0}), + S * k_input.m_element_size); + std::memcpy(&past_v_output.at({b, h, m, 0}), + &v_input.at({b, h, m, 0}), + S * v_input.m_element_size); }); } @@ -82,16 +81,13 @@ void attn_memcpy(const ov::intel_cpu::PlainTensor& k_input, const ov::intel_cpu::PlainTensor& past_k_output, const ov::intel_cpu::PlainTensor& past_v_output) { if (past_k_output.get_precision() == k_input.get_precision()) { - if (past_k_output.get_precision() == ov::element::bf16) { - attn_memcpy_kernel(k_input, v_input, past_k_output, past_v_output); - } else { - assert(past_k_output.get_precision() == ov::element::f16); - attn_memcpy_kernel(k_input, v_input, past_k_output, past_v_output); - } - } else if (past_k_output.get_precision() == ov::element::f16) { + attn_memcpy_kernel(k_input, v_input, past_k_output, past_v_output); + } else if (k_input.get_precision() == ov::element::f32 && past_k_output.get_precision() == ov::element::f16) { attn_memcpy_kernel(k_input, v_input, past_k_output, past_v_output); + } else if (k_input.get_precision() == ov::element::f32 && past_k_output.get_precision() == ov::element::bf16) { + attn_memcpy_kernel(k_input, v_input, past_k_output, past_v_output); } else { - attn_memcpy_kernel(k_input, v_input, past_k_output, past_v_output); + OPENVINO_THROW("unsupport src type: ", k_input.get_precision(), ", dst type: ", past_k_output.get_precision(), " in attn_memcpy"); } } } // namespace XARCH diff --git a/src/plugins/intel_cpu/src/nodes/memory.cpp b/src/plugins/intel_cpu/src/nodes/memory.cpp index 9508cd20cc9..65645c07184 100644 --- a/src/plugins/intel_cpu/src/nodes/memory.cpp +++ b/src/plugins/intel_cpu/src/nodes/memory.cpp @@ -6,6 +6,7 @@ #include #include #include "memory.hpp" +#include "scaled_attn.h" #include "common/cpu_convert.h" #include "common/cpu_memcpy.h" #include "utils/general_utils.h" @@ -21,6 +22,65 @@ namespace ov { namespace intel_cpu { namespace node { +namespace { +class MemoryStub : public IMemory { +public: + MemoryStub(const dnnl::engine& eng, const MemoryDescPtr& pMemDesc) : m_eng(eng), m_pMemDesc(pMemDesc) {} + + bool isAllocated() const noexcept override { + return true; + } + + const MemoryDesc& getDesc() const override { + return *m_pMemDesc; + } + + MemoryDescPtr getDescPtr() const override { + return m_pMemDesc; + } + + void* getData() const override { + OPENVINO_THROW("Unexpected call MemoryStub::getData()"); + } + + size_t getSize() const override { + return 0; + } + + const Shape& getShape() const override { + return m_pMemDesc->getShape(); + } + + const VectorDims& getStaticDims() const override { + return m_pMemDesc->getShape().getStaticDims(); + } + + void redefineDesc(MemoryDescPtr desc) override { + m_pMemDesc = desc; + } + + void load(const IMemory& src, bool ftz = true) const override { + OPENVINO_THROW("Unexpected call MemoryStub::load()"); + } + + MemoryMngrPtr getMemoryMngr() const override { + OPENVINO_THROW("Unexpected call MemoryStub::getMemoryMngr()"); + } + + dnnl::memory getPrimitive() const override { + OPENVINO_THROW("Unexpected call MemoryStub::getPrimitive()"); + } + + void nullify() override { + // nothing to do + } + +private: + dnnl::engine m_eng; + MemoryDescPtr m_pMemDesc; +}; +} // namespace + std::mutex MemoryNodeVirtualEdge::holderMutex; MemoryNode::MemoryNode(const std::shared_ptr& op) { @@ -33,7 +93,7 @@ MemoryNode::MemoryNode(const std::shared_ptr& op) { } } -bool MemoryOutput::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { +bool MemoryOutputBase::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { try { if (!one_of(op->get_type_info(), ov::op::v3::Assign::get_type_info_static(), @@ -47,7 +107,7 @@ bool MemoryOutput::isSupportedOperation(const std::shared_ptr& o return true; } -MemoryOutput::MemoryOutput(const std::shared_ptr& op, const GraphContext::CPtr context) +MemoryOutputBase::MemoryOutputBase(const std::shared_ptr& op, const GraphContext::CPtr context) : Node(op, context, NgraphShapeInferFactory(op, EMPTY_PORT_MASK)) , MemoryNode(op) { std::string errorMessage; if (!isSupportedOperation(op, errorMessage)) { @@ -58,19 +118,34 @@ MemoryOutput::MemoryOutput(const std::shared_ptr& op, const GraphConte } } -MemoryOutput::~MemoryOutput() { +MemoryOutputBase::MemoryOutputBase(const std::string id, + const std::string& name, + const std::string& type, + const Shape& input_shape, + const ov::element::Type& input_prc, + const GraphContext::CPtr context) : + Node(type, name, context), MemoryNode(id) { + isDynamic = input_shape.isDynamic(); + if (isDynamic) { + shapeInference = PassThroughShapeInferFactory().makeShapeInfer(); + } + inputShapes.emplace_back(input_shape); + addOriginalInputPrecision(input_prc); +} + +MemoryOutputBase::~MemoryOutputBase() { if (inputNode) { inputNode->deregisterSibling(this); } MemoryNodeVirtualEdge::remove(this, holder); } -MemoryInputBase& MemoryOutput::getInputNode() { +MemoryInputBase& MemoryOutputBase::getInputNode() { OPENVINO_ASSERT(inputNode, "MemoryOutput ", getName(), " doesn't have sibling input"); return *inputNode; } -void MemoryOutput::getSupportedDescriptors() {} +void MemoryOutputBase::getSupportedDescriptors() {} -void MemoryOutput::initSupportedPrimitiveDescriptors() { +void MemoryOutputBase::initSupportedPrimitiveDescriptors() { if (!supportedPrimitiveDescriptors.empty()) return; @@ -90,7 +165,7 @@ void MemoryOutput::initSupportedPrimitiveDescriptors() { supportedPrimitiveDescriptors.emplace_back(config, impl_desc_type::unknown); } -void MemoryOutput::initOptimalPrimitiveDescriptor() { +void MemoryOutputBase::initOptimalPrimitiveDescriptor() { // Mimic the parent node memory desc to avoid extra reorder auto parentEdge = getParentEdgeAt(0); auto parent = parentEdge->getParent(); @@ -120,6 +195,21 @@ void MemoryOutput::initOptimalPrimitiveDescriptor() { selected_pd->setConfig(config); } +void MemoryOutputBase::registerInputNode(MemoryInputBase* node) { + if (inputNode == node) { return; } + if (inputNode) { inputNode->deregisterSibling(this); } + inputNode = node; + inputNode->registerOutputNode(this); +} + +void MemoryOutputBase::deregisterSibling(MemoryInputBase* node) { + if (node == inputNode) { inputNode = nullptr; } +} + +bool MemoryOutput::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { + return MemoryOutputBase::isSupportedOperation(op, errorMessage); +} + void MemoryOutput::resolveInPlaceEdges(Edge::LOOK look) { if (!(look & Edge::LOOK_DOWN)) { Node::resolveInPlaceEdges(look); @@ -198,17 +288,48 @@ void MemoryOutput::executeDynamicImpl(dnnl::stream strm) { execute(strm); } -void MemoryOutput::registerInputNode(MemoryInputBase* node) { - if (inputNode == node) { return; } - if (inputNode) { inputNode->deregisterSibling(this); } - inputNode = node; - inputNode->registerOutputNode(this); +bool MemoryOutputStub::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { + return MemoryOutputBase::isSupportedOperation(op, errorMessage); } -void MemoryOutput::deregisterSibling(MemoryInputBase* node) { - if (node == inputNode) { inputNode = nullptr; } +void MemoryOutputStub::execute(dnnl::stream strm) { + //nothing to do } +void MemoryOutputStub::executeDynamicImpl(dnnl::stream strm) { + //nothing to do +} + +void MemoryOutputStub::resolveInPlaceEdges(Edge::LOOK look) { + if (!(look & Edge::LOOK_DOWN)) { + Node::resolveInPlaceEdges(look); + return; + } + + auto selected_pd = getSelectedPrimitiveDescriptor(); + OPENVINO_ASSERT(selected_pd, + "MemoryOutput ", + getName(), + " failed getSelectedPrimitiveDescriptor() call, preferable primitive descriptor is not set"); + + auto parentEdge = getParentEdgeAt(0); // always only one parent edge + + OPENVINO_ASSERT(one_of(parentEdge->getStatus(), Edge::Status::Uninitialized, Edge::Status::NotAllocated), + " Unexpected inplace resolve call to an allocated edge: ", parentEdge->name()); + + auto memDesc = selected_pd->getConfig().inConfs.front().getMemDesc(); + // make a fake memory + auto edgeMem = std::make_shared(getEngine(), memDesc); + parentEdge->reuse(edgeMem); +} + +void MemoryOutputStub::assignExtMemory(const MemoryPtr& mem, const MemoryDescPtr& memDesc) { + //nothing to do +} + +bool MemoryOutputStub::isExecutable() const { + return false; +} bool MemoryInputBase::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { try { @@ -260,172 +381,17 @@ MemoryInputBase::MemoryInputBase(const std::string id, // this is their responsibility to link the input/output nodes properly } -void MemoryInputBase::resolveInPlaceEdges(Edge::LOOK look) { - if (!(look & Edge::LOOK_UP)) { - Node::resolveInPlaceEdges(look); - return; - } - - auto selected_pd = getSelectedPrimitiveDescriptor(); - OPENVINO_ASSERT(selected_pd, - "MemoryInput ", - getName(), - " failed getSelectedPrimitiveDescriptor() call, preferable primitive descriptor is not set"); - - auto memDesc = selected_pd->getConfig().outConfs.front().getMemDesc(); - memMngr = std::make_shared(); - - for (auto&& edge : getChildEdgesAtPort(0)) { // always only one child port - OPENVINO_ASSERT(one_of(edge->getStatus(), Edge::Status::Uninitialized, Edge::Status::NotAllocated), - " Unexpected inplace resolve call to an allocated edge: ", edge->name()); - - auto edgeMem = std::make_shared(getEngine(), memDesc, memMngr); - edge->reuse(edgeMem); - } -} - MemoryInputBase::~MemoryInputBase() { if (outputNode) { outputNode->deregisterSibling(this); } MemoryNodeVirtualEdge::remove(this, holder); } -MemoryOutput& MemoryInputBase::getOutputNode() { +MemoryOutputBase& MemoryInputBase::getOutputNode() { OPENVINO_ASSERT(outputNode, "MemoryOutput ", getName(), " doesn't have sibling input"); return *outputNode; } -void MemoryInputBase::assignState(MemStatePtr newState) { - assignedMem = newState->input_mem(); - - if (!getParentEdges().empty() && newState->is_reset_state()) { - isExecutableFlag = true; - } else { - isExecutableFlag = false; - } - - OPENVINO_ASSERT(assignedMem, - "MemoryInput ", - getName(), - " assigned state has null memory ptr"); - - const auto& newDims = assignedMem->getStaticDims(); - MemoryDescPtr internDesc; - if (isDynamicNode()) { - const bool hasZeroDims = std::count(std::begin(newDims), std::end(newDims), 0) > 0; - internDesc = getBaseMemDescAtOutputPort(0)->cloneWithNewDims(newDims, hasZeroDims); - } else { - auto expectedDims = getBaseMemDescAtOutputPort(0)->getShape().getStaticDims(); - OPENVINO_ASSERT(expectedDims == newDims, - "MemoryInput ", - getName(), - " unexpected state shape: ", - vec2str(newDims), - ", while the expected shape: ", - vec2str(expectedDims)); - - internDesc = getBaseMemDescAtOutputPort(0); - } - - OPENVINO_ASSERT(memMngr, - "MemoryInput ", - getName(), - " has uninitialized memory manager."); - - if (internDesc->isCompatible(assignedMem->getDesc())) { - memMngr->setMemMngr(assignedMem->getMemoryMngr()); - } else { - memMngr->reset(); - } - - if (!isExecutableFlag) { - const auto& edges = getChildEdgesAtPort(0); - if (isDynamicNode()) { - for (auto&& edge : edges) { - edge->getMemoryPtr()->redefineDesc(internDesc); - } - } - - auto outMem = edges.front()->getMemoryPtr(); - - if (outMem->getData() != assignedMem->getData()) { - outMem->load(*assignedMem); - } - } - - getOutputNode().assignExtMemory(newState->output_mem(), newState->internal_desc()); -} - -bool MemoryInputBase::needShapeInfer() const { - return isExecutableFlag; -} - -bool MemoryInputBase::isExecutable() const { - return isExecutableFlag && Node::isExecutable(); -} - -void MemoryInputBase::executeDynamicImpl(dnnl::stream strm) { - execute(strm); -} - -void MemoryInputBase::execute(dnnl::stream strm) { - if (!isExecutableFlag) return; - - auto&& src = getParentEdgeAt(0)->getMemory(); - auto&& dst = getChildEdgesAtPort(0).front()->getMemoryPtr(); - dst->load(src); -} - -void MemoryInputBase::registerOutputNode(MemoryOutput* node) { - if (outputNode == node) { return; } - if (outputNode) { outputNode->deregisterSibling(this); } - outputNode = node; - outputNode->registerInputNode(this); -} - -void MemoryInputBase::deregisterSibling(MemoryOutput* node) { - if (node == outputNode) { outputNode = nullptr; } -} - -MemoryNodeVirtualEdge::Holder* MemoryNodeVirtualEdge::registerInput(MemoryInputBase * node) { - std::lock_guard lock{MemoryNodeVirtualEdge::holderMutex}; - // in case of output already registered - auto& holder = MemoryNodeVirtualEdge::getExisted(); - auto sibling = MemoryNodeVirtualEdge::getByName(holder, node->getId()); - if (sibling != nullptr) { - auto outputNode = dynamic_cast(sibling); - OPENVINO_ASSERT(outputNode != nullptr); - node->registerOutputNode(outputNode); - } else { - holder[node->getId()] = node; - } - return &holder; -} - -MemoryNodeVirtualEdge::Holder* MemoryNodeVirtualEdge::registerOutput(MemoryOutput * node) { - std::lock_guard lock{MemoryNodeVirtualEdge::holderMutex}; - // in case of output layer - auto& holder = MemoryNodeVirtualEdge::getExisted(); - auto sibling = MemoryNodeVirtualEdge::getByName(holder, node->getId()); - if (sibling != nullptr) { - auto inputNode = dynamic_cast(sibling); - OPENVINO_ASSERT(inputNode != nullptr); - node->registerInputNode(inputNode); - } else { - holder[node->getId()] = node; - } - return &holder; -} - -void MemoryNodeVirtualEdge::remove(MemoryNode * node, Holder* holder) { - std::lock_guard lock{MemoryNodeVirtualEdge::holderMutex}; - if (nullptr != holder) { - InferenceEngine::details::erase_if(*holder, [&](const Holder::value_type & it){ - return it.second == node; - }); - } -} - -void MemoryInput::initSupportedPrimitiveDescriptors() { +void MemoryInputBase::initSupportedPrimitiveDescriptors() { if (!supportedPrimitiveDescriptors.empty()) return; @@ -456,6 +422,64 @@ void MemoryInput::initSupportedPrimitiveDescriptors() { supportedPrimitiveDescriptors.emplace_back(config, impl_desc_type::unknown); } +void MemoryInputBase::registerOutputNode(MemoryOutputBase* node) { + if (outputNode == node) { return; } + if (outputNode) { outputNode->deregisterSibling(this); } + outputNode = node; + outputNode->registerInputNode(this); +} + +void MemoryInputBase::deregisterSibling(MemoryOutputBase* node) { + if (node == outputNode) { outputNode = nullptr; } +} + +MemoryNodeVirtualEdge::Holder* MemoryNodeVirtualEdge::registerInput(MemoryInputBase * node) { + std::lock_guard lock{MemoryNodeVirtualEdge::holderMutex}; + // in case of output already registered + auto& holder = MemoryNodeVirtualEdge::getExisted(); + auto sibling = MemoryNodeVirtualEdge::getByName(holder, node->getId()); + if (sibling != nullptr) { + auto outputNode = dynamic_cast(sibling); + OPENVINO_ASSERT(outputNode != nullptr); + node->registerOutputNode(outputNode); + } else { + holder[node->getId()] = node; + } + return &holder; +} + +MemoryNodeVirtualEdge::Holder* MemoryNodeVirtualEdge::registerOutput(MemoryOutputBase * node) { + std::lock_guard lock{MemoryNodeVirtualEdge::holderMutex}; + // in case of output layer + auto& holder = MemoryNodeVirtualEdge::getExisted(); + auto sibling = MemoryNodeVirtualEdge::getByName(holder, node->getId()); + if (sibling != nullptr) { + auto inputNode = dynamic_cast(sibling); + OPENVINO_ASSERT(inputNode != nullptr); + node->registerInputNode(inputNode); + } else { + holder[node->getId()] = node; + } + return &holder; +} + +void MemoryNodeVirtualEdge::remove(MemoryNode * node, Holder* holder) { + std::lock_guard lock{MemoryNodeVirtualEdge::holderMutex}; + if (nullptr != holder) { + InferenceEngine::details::erase_if(*holder, [&](const Holder::value_type & it){ + return it.second == node; + }); + } +} + +bool MemoryInput::needShapeInfer() const { + return isExecutableFlag; +} + +bool MemoryInput::isExecutable() const { + return isExecutableFlag && Node::isExecutable(); +} + void MemoryInput::initOptimalPrimitiveDescriptor() { // Mimic the child node memory desc to avoid extra reorder static const Type preferredTypes[] = { @@ -512,6 +536,42 @@ void MemoryInput::initOptimalPrimitiveDescriptor() { selectedPd->setConfig(config); } +void MemoryInput::executeDynamicImpl(dnnl::stream strm) { + execute(strm); +} + +void MemoryInput::execute(dnnl::stream strm) { + if (!isExecutableFlag) return; + + auto&& src = getParentEdgeAt(0)->getMemory(); + auto&& dst = getChildEdgesAtPort(0).front()->getMemoryPtr(); + dst->load(src); +} + +void MemoryInput::resolveInPlaceEdges(Edge::LOOK look) { + if (!(look & Edge::LOOK_UP)) { + Node::resolveInPlaceEdges(look); + return; + } + + auto selected_pd = getSelectedPrimitiveDescriptor(); + OPENVINO_ASSERT(selected_pd, + "MemoryInput ", + getName(), + " failed getSelectedPrimitiveDescriptor() call, preferable primitive descriptor is not set"); + + auto memDesc = selected_pd->getConfig().outConfs.front().getMemDesc(); + memMngr = std::make_shared(); + + for (auto&& edge : getChildEdgesAtPort(0)) { // always only one child port + OPENVINO_ASSERT(one_of(edge->getStatus(), Edge::Status::Uninitialized, Edge::Status::NotAllocated), + " Unexpected inplace resolve call to an allocated edge: ", edge->name()); + + auto edgeMem = std::make_shared(getEngine(), memDesc, memMngr); + edge->reuse(edgeMem); + } +} + MemStatePtr MemoryInput::makeState() const { // assume ov::Tensor is always dense auto original_desc = @@ -534,10 +594,102 @@ MemStatePtr MemoryInput::makeState() const { original_desc); } +void MemoryInput::assignState(MemStatePtr newState) { + assignedMem = newState->input_mem(); + + isExecutableFlag = !getParentEdges().empty() && newState->is_reset_state(); + + OPENVINO_ASSERT(assignedMem, + "MemoryInput ", + getName(), + " assigned state has null memory ptr"); + + const auto& newDims = assignedMem->getStaticDims(); + MemoryDescPtr internDesc; + if (isDynamicNode()) { + const bool hasZeroDims = std::count(std::begin(newDims), std::end(newDims), 0) > 0; + internDesc = getBaseMemDescAtOutputPort(0)->cloneWithNewDims(newDims, hasZeroDims); + } else { + auto expectedDims = getBaseMemDescAtOutputPort(0)->getShape().getStaticDims(); + OPENVINO_ASSERT(expectedDims == newDims, + "MemoryInput ", + getName(), + " unexpected state shape: ", + vec2str(newDims), + ", while the expected shape: ", + vec2str(expectedDims)); + + internDesc = getBaseMemDescAtOutputPort(0); + } + + OPENVINO_ASSERT(memMngr, + "MemoryInput ", + getName(), + " has uninitialized memory manager."); + + if (internDesc->isCompatible(assignedMem->getDesc())) { + memMngr->setMemMngr(assignedMem->getMemoryMngr()); + } else { + memMngr->reset(); + } + + if (!isExecutableFlag) { + const auto& edges = getChildEdgesAtPort(0); + if (isDynamicNode()) { + for (auto&& edge : edges) { + edge->getMemoryPtr()->redefineDesc(internDesc); + } + } + + auto outMem = edges.front()->getMemoryPtr(); + + if (outMem->getData() != assignedMem->getData()) { + outMem->load(*assignedMem); + } + } + + getOutputNode().assignExtMemory(newState->output_mem(), newState->internal_desc()); +} + + bool MemoryInput::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { return MemoryInputBase::isSupportedOperation(op, errorMessage); } +MemoryInputSDPA::MemoryInputSDPA(const std::string id, + const std::string& name, + const std::string& type, + const Shape& output_shape, + const ov::element::Type& output_prc, + const GraphContext::CPtr context, + const ov::optional& input_shape, + const ov::optional& input_prc, + const std::shared_ptr& sdpaNode) : + MemoryInputBase(id, name, type, output_shape, output_prc, context, input_shape, input_prc), m_sdpaNode(sdpaNode) {} + + +bool MemoryInputSDPA::needShapeInfer() const { + return m_needShapeInfer; +} +bool MemoryInputSDPA::isExecutable() const { + return false; +} + +void MemoryInputSDPA::createPrimitive() { + MemoryInputBase::createPrimitive(); + // determine the output node idx + auto memDesc = getBaseMemDescAtOutputPort(0); + auto sdpaNode = m_sdpaNode.lock(); + for (auto&& edge : getChildEdgesAtPort(0)) { // always only one child port + auto node = edge->getChild(); + if (node == sdpaNode) { + m_child_port_idx = edge->getOutputNum(); + break; + } + } + OPENVINO_ASSERT(m_child_port_idx != -1, getName(), " should be connected to SDPA node."); +} + void MemoryInputSDPA::initSupportedPrimitiveDescriptors() { if (!supportedPrimitiveDescriptors.empty()) return; @@ -554,56 +706,45 @@ void MemoryInputSDPA::initSupportedPrimitiveDescriptors() { config.inConfs.push_back(std::move(inPortConfig)); } - auto&& childEdges = getChildEdgesAtPort(0); - auto itr = std::find_if(childEdges.begin(), childEdges.end(), - [](const EdgePtr& edge){ return Type::ScaledDotProductAttention == edge->getChild()->getType(); }); - - OPENVINO_ASSERT(itr != childEdges.end(), "MemoryInputSDPA isn't attached to an SDPA node"); - auto SDPA = (*itr)->getChild(); - auto childPort = (*itr)->getOutputNum(); - - // Since this is a very specialized implementation, lets mimic SDPA precision and set cabd layout - precision = SDPA->getOriginalInputPrecisionAtPort(childPort); - // Just used a place holder here, the actual layout is obtained at initOptimalPrimitiveDescriptor - ArbitraryOrderDescCreator cabdDescCreator({2, 0, 1, 3}); - PortConfig outPortConfig; outPortConfig.inPlace(0); outPortConfig.constant(false); - outPortConfig.setMemDesc(cabdDescCreator.createSharedDesc(precision, shape)); + // layout for fake memory obj, the child sdpa also does not use it + outPortConfig.setMemDesc(descCreators.at(LayoutType::ncsp)->createSharedDesc(precision, shape)); config.outConfs.push_back(std::move(outPortConfig)); supportedPrimitiveDescriptors.emplace_back(config, impl_desc_type::unknown); } void MemoryInputSDPA::initOptimalPrimitiveDescriptor() { - auto&& childEdges = getChildEdgesAtPort(0); - auto itr = std::find_if(childEdges.begin(), childEdges.end(), - [](const EdgePtr& edge){ return Type::ScaledDotProductAttention == edge->getChild()->getType(); }); + Node::initOptimalPrimitiveDescriptor(); +} - OPENVINO_ASSERT(itr != childEdges.end(), "MemoryInputSDPA isn't attached to an SDPA node"); - auto childEdge = *itr; - auto child = childEdge->getChild(); - auto childPd = child->getSelectedPrimitiveDescriptor(); - OPENVINO_ASSERT(childPd, - child->getTypeStr(), " ", - child->getName(), - "failed initOptimalPrimitiveDescriptor() call, preferable primitive descriptor is not set"); +void MemoryInputSDPA::assignState(MemStatePtr newState) { + if (newState->is_reset_state()) { + if (getParentEdges().empty()) { + auto newShape = MemoryDescUtils::makeDummyShape(getBaseMemDescAtOutputPort(0)->getShape(), 0); + redefineOutputMemory({newShape.getStaticDims()}); + m_needShapeInfer = false; + } else { + m_needShapeInfer = true; + } + } else { + auto stateMem = newState->input_mem(); + OPENVINO_ASSERT(stateMem, + "Internal state mem id: ", + newState->get_name(), + " is empty, node name: ", + getName()); - const auto& childConfig = childPd->getConfig(); + redefineOutputMemory({stateMem->getStaticDims()}); + m_needShapeInfer = false; + } - auto selectedPd = getSelectedPrimitiveDescriptor(); - OPENVINO_ASSERT(selectedPd, - "MemoryInputSDPA ", - getName(), - " failed initOptimalPrimitiveDescriptor() call, preferable primitive descriptor is not set"); - - auto config = selectedPd->getConfig(); - // The pyscial layout varies from models, e.g. [LBHS]chatglm, [BHLS]Llama - // The SDPA knows details, so should trust the layout config provided by SPDA - auto newMemDesc = childConfig.inConfs.back().getMemDesc(); - config.outConfs.front().setMemDesc(newMemDesc); - //bypass any checks, we enforce the child descriptor precision - selectedPd->setConfig(config); + auto sdpaNode = m_sdpaNode.lock(); + OPENVINO_ASSERT(sdpaNode); + auto sdpaState = std::dynamic_pointer_cast(newState); + OPENVINO_ASSERT(sdpaState); + sdpaNode->assignState(sdpaState, m_child_port_idx); } MemStatePtr MemoryInputSDPA::makeState() const { @@ -612,7 +753,6 @@ MemStatePtr MemoryInputSDPA::makeState() const { std::make_shared(getOriginalOutputPrecisionAtPort(0), outputShapes.at(0)); auto mem_desc = getBaseMemDescAtOutputPort(0); - const auto& eng = getEngine(); auto state_name = getId(); @@ -622,10 +762,40 @@ MemStatePtr MemoryInputSDPA::makeState() const { state_name = state_name.substr(0, suffix_idx); } - auto internal_memory = - std::make_shared(eng, mem_desc, std::make_shared(make_unique())); + auto node = m_sdpaNode.lock(); + // retrieve the internal precision and axis order from the SDPA node + OPENVINO_ASSERT(node); + auto kv_precision = node->getKVCachePrecision(); + VectorDims order = {0, 1, 2, 3}; + if (!node->getKVCacheOrder().empty()) + order = node->getKVCacheOrder(); - return std::make_shared(state_name, internal_memory, original_desc); + auto internal_desc = ArbitraryOrderDescCreator(order).createSharedDesc(kv_precision, outputShapes.at(0)); + + return std::make_shared(state_name, original_desc, internal_desc); +} + +void MemoryInputSDPA::execute(dnnl::stream strm) { + //nothing to do +} + +void MemoryInputSDPA::executeDynamicImpl(dnnl::stream strm) { + //nothing to do +} + +void MemoryInputSDPA::resolveInPlaceEdges(Edge::LOOK look) { + if (getParentEdgeAt(0)) { + Node::resolveInPlaceEdges(look); + } else { + auto memDesc = getBaseMemDescAtOutputPort(0); + for (auto&& edge : getChildEdgesAtPort(0)) { // always only one child port + OPENVINO_ASSERT(one_of(edge->getStatus(), Edge::Status::Uninitialized, Edge::Status::NotAllocated), + " Unexpected inplace resolve call to an allocated edge: ", edge->name()); + + auto edgeMem = std::make_shared(getEngine(), memDesc); + edge->reuse(edgeMem); + } + } } } // namespace node diff --git a/src/plugins/intel_cpu/src/nodes/memory.hpp b/src/plugins/intel_cpu/src/nodes/memory.hpp index 2d7a4242ac0..7d648e42c63 100644 --- a/src/plugins/intel_cpu/src/nodes/memory.hpp +++ b/src/plugins/intel_cpu/src/nodes/memory.hpp @@ -20,8 +20,9 @@ namespace ov { namespace intel_cpu { namespace node { -class MemoryOutput; +class MemoryOutputBase; class MemoryInputBase; +class ScaledDotProductAttention; class MemoryNode { public: @@ -64,27 +65,31 @@ public: return nullptr; } - static Holder* registerOutput(MemoryOutput * node); + static Holder* registerOutput(MemoryOutputBase * node); static Holder* registerInput(MemoryInputBase * node); static void remove(MemoryNode * node, Holder* holder); static std::mutex holderMutex; }; -class MemoryOutput : public Node, public MemoryNode { +class MemoryOutputBase : public Node, public MemoryNode { public: - MemoryOutput(const std::shared_ptr& op, const GraphContext::CPtr context); - ~MemoryOutput() override; + MemoryOutputBase(const std::shared_ptr& op, const GraphContext::CPtr context); + MemoryOutputBase(const std::string id, + const std::string& name, + const std::string& type, + const Shape& input_shape, + const ov::element::Type& input_prc, + const GraphContext::CPtr context); + + ~MemoryOutputBase() override; static bool isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept; void getSupportedDescriptors() override; void initSupportedPrimitiveDescriptors() override; void initOptimalPrimitiveDescriptor() override; void createPrimitive() override {} - void execute(dnnl::stream strm) override; - void executeDynamicImpl(dnnl::stream strm) override; bool created() const override { return getType() == Type::MemoryOutput; } - void resolveInPlaceEdges(Edge::LOOK look) override; void registerInputNode(MemoryInputBase* node); void deregisterSibling(MemoryInputBase* node); @@ -92,9 +97,9 @@ public: bool needShapeInfer() const override { return false; } bool needPrepareParams() const override { return false; } - void assignExtMemory(const MemoryPtr& mem, const MemoryDescPtr& memDesc); + virtual void assignExtMemory(const MemoryPtr& mem, const MemoryDescPtr& memDesc) = 0; -private: +protected: MemoryInputBase& getInputNode(); private: @@ -102,12 +107,41 @@ private: * @brief keeps reference to input sibling node */ MemoryInputBase* inputNode = nullptr; + MemoryNodeVirtualEdge::Holder* holder = nullptr; +}; + +class MemoryOutput : public MemoryOutputBase { +public: + using MemoryOutputBase::MemoryOutputBase; + static bool isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept; + + void execute(dnnl::stream strm) override; + void executeDynamicImpl(dnnl::stream strm) override; + + void resolveInPlaceEdges(Edge::LOOK look) override; + + void assignExtMemory(const MemoryPtr& mem, const MemoryDescPtr& memDesc) override; + +private: MemoryPtr assignedMem = nullptr; MemoryDescPtr extMemDesc = nullptr; // used for resize - MemoryNodeVirtualEdge::Holder* holder = nullptr; ProxyMemoryMngrPtr memMngr = nullptr; }; +class MemoryOutputStub : public MemoryOutputBase { +public: + using MemoryOutputBase::MemoryOutputBase; + static bool isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept; + + void execute(dnnl::stream strm) override; + void executeDynamicImpl(dnnl::stream strm) override; + + void resolveInPlaceEdges(Edge::LOOK look) override; + + void assignExtMemory(const MemoryPtr& mem, const MemoryDescPtr& memDesc) override; + bool isExecutable() const override; +}; + class MemoryInputBase : public Input, public MemoryStateNode { public: MemoryInputBase(const std::shared_ptr& op, const GraphContext::CPtr context); @@ -127,29 +161,19 @@ public: return getType() == Type::MemoryInput; } - bool needShapeInfer() const override; - bool isExecutable() const override; - void execute(dnnl::stream strm) override; - void executeDynamicImpl(dnnl::stream strm) override; + void initSupportedPrimitiveDescriptors() override; - void resolveInPlaceEdges(Edge::LOOK look) override; + void registerOutputNode(MemoryOutputBase* node); + void deregisterSibling(MemoryOutputBase* node); - void registerOutputNode(MemoryOutput* node); - void deregisterSibling(MemoryOutput* node); - - // May be extracted to some interface when necessary - void assignState(MemStatePtr newState) override; - MemoryOutput& getOutputNode(); + MemoryOutputBase& getOutputNode(); private: /** * @brief keeps reference to output sibling node */ - MemoryOutput* outputNode = nullptr; - MemoryPtr assignedMem = nullptr; + MemoryOutputBase* outputNode = nullptr; MemoryNodeVirtualEdge::Holder* holder = nullptr; - ProxyMemoryMngrPtr memMngr = nullptr; - bool isExecutableFlag = true; }; class MemoryInput : public MemoryInputBase { @@ -157,21 +181,56 @@ public: using MemoryInputBase::MemoryInputBase; static bool isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept; - void initSupportedPrimitiveDescriptors() override; + bool needShapeInfer() const override; + bool isExecutable() const override; void initOptimalPrimitiveDescriptor() override; + void execute(dnnl::stream strm) override; + void executeDynamicImpl(dnnl::stream strm) override; + void resolveInPlaceEdges(Edge::LOOK look) override; + + void assignState(MemStatePtr newState) override; MemStatePtr makeState() const override; + +private: + bool isExecutableFlag = true; + ProxyMemoryMngrPtr memMngr = nullptr; + MemoryPtr assignedMem = nullptr; }; class MemoryInputSDPA : public MemoryInputBase { public: - using MemoryInputBase::MemoryInputBase; + MemoryInputSDPA(const std::string id, + const std::string& name, + const std::string& type, + const Shape& output_shape, + const ov::element::Type& output_prc, + const GraphContext::CPtr context, + const ov::optional& input_shape, + const ov::optional& input_prc, + const std::shared_ptr& sdpaNode); + static bool isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept; + bool needShapeInfer() const override; + bool isExecutable() const override; + + void createPrimitive() override; void initSupportedPrimitiveDescriptors() override; void initOptimalPrimitiveDescriptor() override; + void execute(dnnl::stream strm) override; + void executeDynamicImpl(dnnl::stream strm) override; + + void resolveInPlaceEdges(Edge::LOOK look) override; + + void assignState(MemStatePtr newState) override; MemStatePtr makeState() const override; + +private: + std::weak_ptr m_sdpaNode; + int m_child_port_idx = -1; + bool m_needShapeInfer = false; }; } // namespace node } // namespace intel_cpu diff --git a/src/plugins/intel_cpu/src/nodes/scaled_attn.cpp b/src/plugins/intel_cpu/src/nodes/scaled_attn.cpp index fad65cbbd7e..3dc40ff00c8 100644 --- a/src/plugins/intel_cpu/src/nodes/scaled_attn.cpp +++ b/src/plugins/intel_cpu/src/nodes/scaled_attn.cpp @@ -18,9 +18,10 @@ #include "openvino/core/parallel.hpp" #include "memory_desc/cpu_memory_desc_utils.h" #include "memory_desc/dnnl_blocked_memory_desc.h" -#include "utils/plain_tensor.hpp" #include #include "common/arbitrary_order_desc_creator.h" +#include +#include "openvino/util/common_util.hpp" #ifdef OV_CPU_WITH_MLAS # include "mlas/sgemm.hpp" @@ -33,12 +34,33 @@ using namespace InferenceEngine; using namespace InferenceEngine::Extensions::Cpu::XARCH; +using namespace dnnl::impl; using namespace dnnl::impl::cpu::x64; namespace ov { namespace intel_cpu { namespace node { +struct ScaledDotProductAttentionKey { + ov::element::Type rtPrecision; + + size_t hash() const; + bool operator==(const ScaledDotProductAttentionKey& rhs) const; +}; + +size_t ScaledDotProductAttentionKey::hash() const { + size_t seed = 0; + seed = hash_combine(seed, rtPrecision.hash()); + + return seed; +} + +bool ScaledDotProductAttentionKey::operator==(const ScaledDotProductAttentionKey& rhs) const { + auto retVal = rtPrecision == rhs.rtPrecision; + + return retVal; +} + // default implementation: reference template struct MHAKernel { @@ -184,7 +206,17 @@ struct MHAKernel { using tag = dnnl::memory::format_tag; using dt = dnnl::memory::data_type; - void prepare_prim(dnnl::stream strm, size_t B, size_t H, size_t Hk, size_t q_len, size_t kv_len, size_t S, bool has_out_transpose) { + void prepare_prim(dnnl::stream strm, + PlainTensor& query, + PlainTensor& present_key, + PlainTensor& present_value, + size_t B, + size_t H, + size_t Hk, + size_t q_len, + size_t kv_len, + size_t S, + bool has_out_transpose) { auto make_dnnl_dims = [](const std::vector& dims) { dnnl::memory::dims dnnl_dims(dims.size()); for (size_t i = 0; i < dims.size(); i++) @@ -192,8 +224,8 @@ struct MHAKernel { return dnnl_dims; }; auto qkv_dt = precision_of::value == ov::element::f32 ? dt::f32 : dt::bf16; - dnnl::memory::desc cur_q_md(make_dnnl_dims({B, H, q_len, S}), qkv_dt, tag::abcd); - dnnl::memory::desc cur_k_md(make_dnnl_dims({B, Hk, kv_len, S}), qkv_dt, tag::abcd); + dnnl::memory::desc cur_q_md(make_dnnl_dims({B, H, q_len, S}), qkv_dt, query.get_strides()); + dnnl::memory::desc cur_k_md(make_dnnl_dims({B, Hk, kv_len, S}), qkv_dt, present_key.get_strides()); if (cur_q_md == q_md && cur_k_md == k_md) return; @@ -205,7 +237,7 @@ struct MHAKernel { qk_prim = dnnl::matmul(qk_pd); weight_md = dnnl::memory::desc(make_dnnl_dims({B, H, q_len, kv_len}), qkv_dt, tag::abcd); - v_md = dnnl::memory::desc(make_dnnl_dims({B, Hk, kv_len, S}), qkv_dt, tag::abcd); + v_md = dnnl::memory::desc(make_dnnl_dims({B, Hk, kv_len, S}), qkv_dt, present_value.get_strides()); out_md = dnnl::memory::desc(make_dnnl_dims({B, H, q_len, S}), qkv_dt, tag::abcd); if (has_out_transpose) out_md = out_md.permute_axes({0, 2, 1, 3}); @@ -266,7 +298,7 @@ struct MHAKernel { if (d_scale == 0.0f) d_scale = 1.0f / sqrt(head_size); - prepare_prim(strm, B, H, Hk, q_len, kv_len, head_size, has_out_transpose); + prepare_prim(strm, query, present_key, present_value, B, H, Hk, q_len, kv_len, head_size, has_out_transpose); exec_qk(strm, query, present_key); PlainTensor score; @@ -479,20 +511,12 @@ struct MHASingleToken { template struct ScaledDotProductAttention::AttentionExecutor : public ScaledDotProductAttention::Executor { - PlainTensor q_input; // f32[B, H, L1, S] - PlainTensor k_input; // f32[B, H|1, L1, S] / [B, H|1, L0+L1, S] - PlainTensor v_input; // f32[B, H|1, L1, S] / [B, H|1, L0+L1, S] - PlainTensor beam_table; // i32[B, max_kvLen] PlainTensor attn_buf; // f32[[B|1],[H|1], L1|1, L0+L1] - float scale_input = 0.0f; MHAKernel kernel; MHASingleToken kernel_single_token; - size_t B, H, L1, L0, S; - - Config config; - AttentionExecutor(const Config& _config) : attn_buf(true), config(_config) {} + AttentionExecutor() : attn_buf(true) {} void prepare_attn_mask(MemoryPtr attn_input) { attn_buf.resize(attn_input->getStaticDims()); @@ -501,54 +525,28 @@ struct ScaledDotProductAttention::AttentionExecutor : public ScaledDotProductAtt attn_buf.data()[i] = p[i] ? 0.0f : -FLT_MAX; } - void concat_pastkv(const std::vector& inputs, - const std::vector& outputs, - const PlainTensor& k_input, - const PlainTensor& v_input, - PlainTensor& past_k_output, - PlainTensor& past_v_output) { - if (config.config.fuse_concat) { - k_input.assert_dims({B, 0, L1, S}, true); - v_input.assert_dims({B, 0, L1, S}, true); - auto past_k_idx = inputs.size() - 2; - auto past_k_mem = inputs[past_k_idx + 0]; - const auto& permute_axes = config.config.permute_axes; - L0 = permute_axes.empty() ? past_k_mem->getStaticDims()[2] : past_k_mem->getStaticDims()[permute_axes[2]]; - // [B, H, L0, S] - past_k_output.reset(outputs[1]); - past_v_output.reset(outputs[2]); - if (!permute_axes.empty()) { - // [L, B, H, S] -> [B, H, L, S] - past_k_output = past_k_output.permute(permute_axes); - past_v_output = past_v_output.permute(permute_axes); - } - attn_memcpy(k_input, v_input, past_k_output.slice(2, L0, L0 + L1), past_v_output.slice(2, L0, L0 + L1)); - if (!config.is_concat_inplaced) { - PlainTensor past_k_input, past_v_input; - past_k_input.reset(past_k_mem); - past_v_input.reset(inputs[past_k_idx + 1]); - attn_memcpy(past_k_input, past_v_input, past_k_output, past_v_output); - } - } else { - // k,v inputs are already concatenated - L0 = k_input.size(2) - L1; - k_input.assert_dims({B, 0, L0 + L1, S}, true); - v_input.assert_dims({B, 0, L0 + L1, S}, true); - past_k_output = k_input; - past_v_output = v_input; - } - } - - void execute(dnnl::stream strm, const std::vector& inputs, const std::vector& outputs) override { + void execute(dnnl::stream strm, const Config& config, const std::vector& inputs, const MemoryPtr output, + const MemoryPtr presentk_input, const MemoryPtr presentv_input, const MemoryPtr beam_input) override { bool has_out_transpose = config.config.output_BLHxS; bool fuse_causal_attn = config.config.fuse_causal_attn; bool is_causal = config.config.is_causal; - const bool fuse_concat = config.config.fuse_concat; - auto input_num = inputs.size() - (fuse_concat ? 2 : 0); + bool fuse_concat = config.config.fuse_concat; + auto input_num = inputs.size(); + PlainTensor present_key, present_value; + PlainTensor q_input; // f32[B, H, L1, S] + PlainTensor k_input; // f32[B, H|1, L1, S] / [B, H|1, L0+L1, S] + PlainTensor v_input; // f32[B, H|1, L1, S] / [B, H|1, L0+L1, S] + PlainTensor beam_table; // i32[B, max_kvLen] + float scale_input = 0.0f; + size_t B, L1, L0, S; q_input.reset(inputs[0]); k_input.reset(inputs[1]); v_input.reset(inputs[2]); + present_key.reset(presentk_input); + present_value.reset(presentv_input); + if (beam_input) + beam_table.reset(beam_input); PlainTensor attn_mask; if (input_num > 3) { // attn_mask @@ -567,20 +565,32 @@ struct ScaledDotProductAttention::AttentionExecutor : public ScaledDotProductAtt // q: [B, H, L1, S] const auto & permute_axes = config.config.permute_axes; - - PlainTensor present_key, present_value; if (!permute_axes.empty()) { q_input = q_input.permute(permute_axes); k_input = k_input.permute(permute_axes); v_input = v_input.permute(permute_axes); + present_key = present_key.permute(permute_axes); + present_value = present_value.permute(permute_axes); } B = q_input.size(0); - H = q_input.size(1); L1 = q_input.size(2); - S = q_input.size(-1); - concat_pastkv(inputs, outputs, k_input, v_input, present_key, present_value); + S = q_input.size(3); + L0 = present_key.size(2) - L1; + auto Hk = k_input.size(1); - ov::intel_cpu::PlainTensor output_emb(outputs[0]); + if (fuse_concat) { + k_input.assert_dims({B, Hk, L1, S}); + v_input.assert_dims({B, Hk, L1, S}); + } else { + k_input.assert_dims({B, Hk, L0 + L1, S}); + v_input.assert_dims({B, Hk, L0 + L1, S}); + } + present_key.assert_dims({B, Hk, L0 + L1, S}); + present_value.assert_dims({B, Hk, L0 + L1, S}); + if (beam_table) + beam_table.assert_dims({B, L0 + L1}); + + ov::intel_cpu::PlainTensor output_emb(output); bool auto_causal; bool use_attn_mask; @@ -611,7 +621,9 @@ struct ScaledDotProductAttention::AttentionExecutor : public ScaledDotProductAtt } } - if (L1 > 1) { + // second token, or first token with pastkv fusing + bool use_one_token = L1 == 1 || (fuse_concat && L0 > 0); + if (!use_one_token) { // multi-token version kernel(strm, q_input, k_input, v_input, {}, use_attn_mask ? attn_mask : PlainTensor(), output_emb, has_out_transpose, auto_causal, scale_input); @@ -628,7 +640,7 @@ struct ScaledDotProductAttention::AttentionExecutor : public ScaledDotProductAtt }; ScaledDotProductAttention::ScaledDotProductAttention(const std::shared_ptr& op, const GraphContext::CPtr context) - : Node(op, context, NgraphShapeInferFactory(op, EMPTY_PORT_MASK)) { + : Node(op, context, NgraphShapeInferFactory(op, EMPTY_PORT_MASK)), m_tmp_reorder(true) { std::string errorMessage; if (!isSupportedOperation(op, errorMessage)) { OPENVINO_THROW("CPU: " + errorMessage); @@ -646,23 +658,8 @@ ScaledDotProductAttention::ScaledDotProductAttention(const std::shared_ptrpermute[1,2,0,3] ->[B,H,L,S] - // The actual index of B is permute[0], H is permute[1], L is permute[2], S is permute[3] - layoutDescCreator = ArbitraryOrderDescCreator({static_cast(permute_axes[2]), - static_cast(permute_axes[0]), - static_cast(permute_axes[1]), - static_cast(permute_axes[3])}); - } - config.inConfs[orginSDPInputNumber + 0].setMemDesc(layoutDescCreator.createSharedDesc( - kvCachePrecision, getInputShapeAtPort(orginSDPInputNumber + 0))); - config.inConfs[orginSDPInputNumber + 1].setMemDesc(layoutDescCreator.createSharedDesc( - kvCachePrecision, getInputShapeAtPort(orginSDPInputNumber + 1))); + // beam_idx + config.inConfs[orginSDPInputNumber + 0].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( + ov::element::i32, getInputShapeAtPort(orginSDPInputNumber + 0))); - config.outConfs[1].setMemDesc(layoutDescCreator.createSharedDesc( - kvCachePrecision, getOutputShapeAtPort(1))); - config.outConfs[1].inPlace(orginSDPInputNumber + 0); - config.outConfs[2].setMemDesc(layoutDescCreator.createSharedDesc( - kvCachePrecision, getOutputShapeAtPort(2))); - config.outConfs[2].inPlace(orginSDPInputNumber + 1); + // Since the InputMemory nodes are simple proxy for the state memory as well as the init subgraph memory, + // it doesn't make sense to set the real KV cache precision, since we don't need any precision conversions + // provided by the common graph logic. We set precisions equal to the precisions of the state nodes to avoid + // reorder insertion in between MemoryInputSDPA and SDPA nodes. + + auto past_k_input_mem_precision = getParentEdgeAt(orginSDPInputNumber + 1)->getParent()->getOriginalOutputPrecisionAtPort(0); + // pastk + config.inConfs[orginSDPInputNumber + 1].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( + past_k_input_mem_precision, getInputShapeAtPort(orginSDPInputNumber + 1))); + + auto past_v_input_mem_precision = getParentEdgeAt(orginSDPInputNumber + 2)->getParent()->getOriginalOutputPrecisionAtPort(0); + // pastv + config.inConfs[orginSDPInputNumber + 2].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( + past_v_input_mem_precision, getInputShapeAtPort(orginSDPInputNumber + 2))); + + config.outConfs[1].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( + past_k_input_mem_precision, getOutputShapeAtPort(1))); + config.outConfs[1].inPlace(-1); + config.outConfs[2].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( + past_v_input_mem_precision, getOutputShapeAtPort(2))); + config.outConfs[2].inPlace(-1); } config.outConfs[0].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( rtPrecision, getOutputShapeAtPort(0))); supportedPrimitiveDescriptors.emplace_back(config, impl_desc_type::ref_any); - // may fallback to abcd without inplace - if (m_config.config.fuse_concat) { - config.inConfs[orginSDPInputNumber + 0].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( - kvCachePrecision, getInputShapeAtPort(orginSDPInputNumber + 0))); - config.inConfs[orginSDPInputNumber + 1].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( - kvCachePrecision, getInputShapeAtPort(orginSDPInputNumber + 1))); - config.outConfs[1].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( - kvCachePrecision, getOutputShapeAtPort(1))); - config.outConfs[1].inPlace(-1); - config.outConfs[2].setMemDesc(creatorsMap.at(LayoutType::ncsp)->createSharedDesc( - kvCachePrecision, getOutputShapeAtPort(2))); - config.outConfs[2].inPlace(-1); - supportedPrimitiveDescriptors.emplace_back(config, impl_desc_type::ref_any); - } } void ScaledDotProductAttention::createPrimitive() { @@ -740,32 +727,51 @@ void ScaledDotProductAttention::createPrimitive() { auto desc = getSelectedPrimitiveDescriptor(); if (desc == nullptr) OPENVINO_THROW("has unidentified preferable primitive descriptor"); - - m_config.is_concat_inplaced = desc->getConfig().outConfs[1].inPlace() >= 0; } + auto rtPrecision = getRuntimePrecision(); - if (rtPrecision == ov::element::bf16) { - m_executor = std::make_shared>(m_config); - } else { - // only support bf16/f32 - rtPrecision = ov::element::f32; -#ifdef OV_CPU_WITH_MLAS - m_executor = std::make_shared>(m_config); -#else - m_executor = std::make_shared>(m_config); -#endif - } + ScaledDotProductAttentionKey key = {rtPrecision}; + + auto builder = [&](const ScaledDotProductAttentionKey& key) -> std::shared_ptr { + std::shared_ptr executor; + if (rtPrecision == ov::element::bf16) { + executor = std::make_shared>(); + } else { + #ifdef OV_CPU_WITH_MLAS + executor = std::make_shared>(); + #else + executor = std::make_shared>(); + #endif + } + return executor; + }; + + auto cache = context->getParamsCache(); + auto result = cache->getOrCreate(key, builder); + m_executor = result.first; } void ScaledDotProductAttention::execute(dnnl::stream strm) { - std::vector inputs(getParentEdges().size()), outputs(getChildEdges().size()); - for (size_t i = 0; i < inputs.size(); i++) { + auto orginSDPInputNumber = getOriginalInputsNumber() - (m_config.config.fuse_concat ? 3 : 0); + std::vector inputs(orginSDPInputNumber); + auto output = getChildEdgeAt(0)->getMemoryPtr(); + MemoryPtr presentk_input, presentv_input, beam_input; + for (size_t i = 0; i < orginSDPInputNumber; i++) { inputs[i] = getParentEdgeAt(i)->getMemoryPtr(); } - for (size_t i = 0; i < outputs.size(); i++) { - outputs[i] = getChildEdgeAt(i)->getMemoryPtr(); + + if (m_config.config.fuse_concat) { + // initialization will be also completed in this func + gatherConcatPastkv(inputs[1], inputs[2], getParentEdgeAt(orginSDPInputNumber)->getMemoryPtr()); + + presentk_input = m_k_state->internal_state_mem(); + presentv_input = m_v_state->internal_state_mem(); + beam_input = m_k_state->hidden_state_mem(); + } else { + presentk_input = inputs[1]; + presentv_input = inputs[2]; } - m_executor->execute(strm, inputs, outputs); + m_executor->execute(strm, m_config, inputs, output, presentk_input, presentv_input, beam_input); } bool ScaledDotProductAttention::isSupportedOperation(const std::shared_ptr& op, std::string& errorMessage) noexcept { @@ -785,7 +791,7 @@ bool ScaledDotProductAttention::isSupportedOperation(const std::shared_ptr(op); if (node) { if (node->get_config().fuse_concat) { - orgSDPAInput -= 2; + orgSDPAInput -= 3; } } if (orgSDPAInput > 3) { @@ -806,6 +812,272 @@ bool ScaledDotProductAttention::isSupportedOperation(const std::shared_ptr& state, int idx) { + auto inputNumber = getOriginalInputsNumber(); + if (inputNumber - 2 == static_cast(idx)) { + m_k_state = state; + } else if (inputNumber - 1 == static_cast(idx)) { + m_v_state = state; + } else { + OPENVINO_THROW( + "Unexpected idx ", idx , " for a state in a node with type: ", getTypeStr(), " and name ", getName()); + } +} + +void ScaledDotProductAttention::gatherConcatPastkv(const MemoryPtr& mem_cur_k, const MemoryPtr& mem_cur_v, const MemoryPtr& mem_beam_idx) { + PlainTensor cur_k; + cur_k.reset(mem_cur_k); + if (!m_config.config.permute_axes.empty()) + cur_k = cur_k.permute(m_config.config.permute_axes); + + updateBeamTable(mem_beam_idx, cur_k.size(2)); + updatePastkv(mem_cur_k, mem_cur_v); +} + +// Update beam table using beam_idx. For first token, beam table is like [[0, 0, 0, ...], [1, 1, 1, ...], ...], +// for second token, beam table is updated using gather(beam_table, beam_idx) then appending [0, 1, 2, ...] to the end for itself. +void ScaledDotProductAttention::updateBeamTable(const MemoryPtr& mem_beam_idx, size_t L1) { + std::vector order = {0, 1, 2, 3}; + if (!m_config.config.permute_axes.empty()) { + order = m_config.config.permute_axes; + } + PlainTensor beam_idx, beam_table_k, beam_table_v; + auto hidden_state_k = m_k_state->hidden_state_mem(); + auto hidden_state_v = m_v_state->hidden_state_mem(); + beam_idx.reset(mem_beam_idx); + + auto B = beam_idx.size(0); + auto is_reset = m_k_state->is_reset_state() || m_v_state->is_reset_state(); + auto inputNumber = getOriginalInputsNumber(); + auto&& v_dims = getParentEdgeAt(inputNumber - 1)->getMemory().getStaticDims(); + size_t L0 = v_dims.at(order[2]); + auto B_state = v_dims.at(order[0]); + OPENVINO_ASSERT(m_k_state->is_reset_state() == m_v_state->is_reset_state(), + "KV state must be reset simultaneously, please also reset state for ", + (m_k_state->is_reset_state() ? m_v_state->get_name() : m_k_state->get_name())); + OPENVINO_ASSERT(B == B_state, "beam idx batch: ", B, " is not equal to batch of state: ", B_state); + OPENVINO_ASSERT(B * (L0 + L1) > 0, "B or (L0+L1) is zero, B: ", B, ", L0: ", L0, ", L1: ", L1); + // resize buffer + if (B * (L0 + L1) > m_k_state->hidden_state_max_size()) { + auto mem_desc = std::make_shared(ov::element::i32, Shape{B, (L0 + L1) * 2}); + + auto new_hidden_state_k = std::make_shared(getEngine(), mem_desc); + auto new_hidden_state_v = std::make_shared(getEngine(), mem_desc); + PlainTensor new_beam_table_k, new_beam_table_v; + new_beam_table_k.reset(new_hidden_state_k); + new_beam_table_v.reset(new_hidden_state_v); + if (L0 > 0 && !is_reset) { + beam_table_k.reset(hidden_state_k); + beam_table_v.reset(hidden_state_v); + for (size_t b = 0; b < B; b++) { + std::memcpy(&new_beam_table_k.at({b}), &beam_table_k.at({b}), sizeof(int32_t) * L0); + std::memcpy(&new_beam_table_v.at({b}), &beam_table_v.at({b}), sizeof(int32_t) * L0); + } + } + m_k_state->assign_hidden_state(new_hidden_state_k); + m_v_state->assign_hidden_state(new_hidden_state_v); + m_k_state->assign_hidden_state_max_size(B * (L0 + L1) * 2); + m_v_state->assign_hidden_state_max_size(B * (L0 + L1) * 2); + hidden_state_k = new_hidden_state_k; + hidden_state_v = new_hidden_state_v; + beam_table_k = new_beam_table_k; + beam_table_v = new_beam_table_v; + } + std::vector new_shape{B, (L0 + L1)}; + auto mem_desc = std::make_shared(ov::element::i32, + Shape(new_shape), + new_shape, + VectorDims{0, 1}, + 0, + VectorDims{}, + hidden_state_k->getDescWithType()->getStrides()); + hidden_state_k->redefineDesc(mem_desc); + hidden_state_v->redefineDesc(mem_desc); + + if (!beam_table_k) { + beam_table_k.reset(hidden_state_k); + beam_table_v.reset(hidden_state_v); + } + + // first token + if (L0 == 0 || is_reset) { + for (size_t b = 0; b < B; b++) { + for (size_t l = 0; l < L0 + L1; l++) { + beam_table_k.at({b, l}) = b; + beam_table_v.at({b, l}) = b; + } + } + return; + } + + // beam order is like [0, 1, 2,...] + bool no_reorder = true; + for (size_t i = 0; i < B; i++) { + if (beam_idx.data()[i] != static_cast(i)) { + no_reorder = false; + break; + } + } + + // reorder + if (!no_reorder) { + m_tmp_reorder.resize({B, L0}); + for (size_t i = 0; i < B; i++) { + std::memcpy(&m_tmp_reorder.at({i}), + &beam_table_k.at({i}), + sizeof(int32_t) * L0); + } + auto* table = beam_idx.data(); + // beam table is same for both k,v state + for (size_t i = 0; i < B; i++) { + std::memcpy(&beam_table_k.at({i}), + &m_tmp_reorder.at({static_cast(table[i])}), + sizeof(int32_t) * L0); + std::memcpy(&beam_table_v.at({i}), + &m_tmp_reorder.at({static_cast(table[i])}), + sizeof(int32_t) * L0); + } + } + // second token itself + for (size_t i = 0; i < B; i++) { + for (size_t j = 0; j < L1; j++) { + beam_table_k.at({i, L0 + j}) = i; + beam_table_v.at({i, L0 + j}) = i; + } + } +} + +// Update pastkv using cur_k, cur_v, simply append cur_k, cur_v to the end of pastkv in the state. +void ScaledDotProductAttention::updatePastkv(const MemoryPtr& mem_cur_k, const MemoryPtr& mem_cur_v) { + std::vector order = {0, 1, 2, 3}; + if (!m_config.config.permute_axes.empty()) { + order = m_config.config.permute_axes; + } + PlainTensor cur_k, past_k; + PlainTensor cur_v, past_v; + cur_k.reset(mem_cur_k); + cur_v.reset(mem_cur_v); + cur_k = cur_k.permute(order); + cur_v = cur_v.permute(order); + auto B = cur_k.size(0); + auto H = cur_k.size(1); + auto L1 = cur_k.size(2); + auto S = cur_k.size(3); + auto reverse = [&order] (const std::vector& cur) { + std::vector result(cur.size()); + for (size_t i = 0; i < cur.size(); i++) { + result[order[i]] = cur[i]; + } + return result; + }; + auto internal_mem_k = m_k_state->internal_state_mem(); + auto internal_mem_v = m_v_state->internal_state_mem(); + + auto is_reset = m_k_state->is_reset_state(); + auto inputNumber = getOriginalInputsNumber(); + auto&& v_dims = getParentEdgeAt(inputNumber - 1)->getMemory().getStaticDims(); + size_t L0 = v_dims.at(order[2]); + auto B_state = v_dims.at(order[0]); + OPENVINO_ASSERT(B == B_state, "pastkv batch: ", B, " is not equal to batch of state: ", B_state); + OPENVINO_ASSERT(B * (L0 + L1) > 0, "B or (L0+L1) is zero, B: ", B, ", L0: ", L0, ", L1: ", L1); + // resize buffer + if (B * H * (L0 + L1) * S > m_k_state->internal_state_max_size()) { + auto new_shape = {B, H, (L0 + L1) * 2, S}; + auto mem_desc = std::make_shared(m_kvcache_precision, + Shape(reverse(new_shape)), + new_shape, + order); + + auto new_internal_mem_k = std::make_shared(getEngine(), mem_desc); + auto new_internal_mem_v = std::make_shared(getEngine(), mem_desc); + + PlainTensor new_pastk, new_pastv; + new_pastk.reset(new_internal_mem_k); + new_pastv.reset(new_internal_mem_v); + new_pastk = new_pastk.permute(order); + new_pastv = new_pastv.permute(order); + if (L0 > 0 && !is_reset) { + past_k.reset(internal_mem_k); + past_v.reset(internal_mem_v); + past_k = past_k.permute(order); + past_v = past_v.permute(order); + attn_memcpy(past_k, past_v, new_pastk, new_pastv); + } + internal_mem_k = new_internal_mem_k; + internal_mem_v = new_internal_mem_v; + past_k = new_pastk; + past_v = new_pastv; + m_k_state->assign_internal_state(new_internal_mem_k); + m_v_state->assign_internal_state(new_internal_mem_v); + m_k_state->assign_internal_state_max_size(B * H * (L0 + L1) * 2 * S); + m_v_state->assign_internal_state_max_size(B * H * (L0 + L1) * 2 * S); + } + auto new_shape = {B, H, (L0 + L1), S}; + auto mem_desc = std::make_shared(m_kvcache_precision, + Shape(reverse(new_shape)), + new_shape, + order, + 0, + VectorDims{}, + internal_mem_k->getDescWithType()->getStrides()); + internal_mem_k->redefineDesc(mem_desc); + internal_mem_v->redefineDesc(mem_desc); + + if (!past_k) { + past_k.reset(internal_mem_k); + past_v.reset(internal_mem_v); + past_k = past_k.permute(order); + past_v = past_v.permute(order); + } + if (L0 > 0 && is_reset) { + auto inputNumber = getOriginalInputsNumber(); + auto k_mem = getParentEdgeAt(inputNumber - 2)->getMemoryPtr(); + auto v_mem = getParentEdgeAt(inputNumber - 1)->getMemoryPtr(); + auto&& k_shape = k_mem->getShape(); + auto&& v_shape = v_mem->getShape(); + if (!k_shape.hasZeroDims() && !v_shape.hasZeroDims()) { + PlainTensor init_k, init_v; + init_k.reset(k_mem); + init_v.reset(v_mem); + init_k = init_k.permute(order); + init_v = init_v.permute(order); + attn_memcpy(init_k, init_v, past_k, past_v); + } + } + + attn_memcpy(cur_k, cur_v, past_k.slice(2, L0, L0 + L1), past_v.slice(2, L0, L0 + L1)); +} + +ov::element::Type ScaledDotProductAttention::getKVCachePrecision() { + if (m_kvcache_precision != ov::element::undefined) + return m_kvcache_precision; + auto rtPrecision = getRuntimePrecision(); + bool enableKVCacheFP16 = m_config.config.fuse_concat && mayiuse(cpu_isa_t::avx2) && rtPrecision != ov::element::bf16; + m_kvcache_precision = enableKVCacheFP16 ? ov::element::f16 : rtPrecision; + + return m_kvcache_precision; +} + +ov::element::Type ScaledDotProductAttention::getRuntimePrecision() const { + auto rtPrecision = getOriginalInputPrecisionAtPort(0); + // only support bf16 and f32 + if (rtPrecision != ov::element::bf16 && rtPrecision != ov::element::f32) + rtPrecision = ov::element::f32; + + size_t H_idx = 1; + if (!m_config.config.permute_axes.empty()) { + H_idx = m_config.config.permute_axes[1]; + } + const auto& qDims = getInputShapeAtPort(0).getDims(); + const auto& kDims = getInputShapeAtPort(1).getDims(); + // if multi-query, enforce fp32 TODO: support BF16 + if (qDims[H_idx] != kDims[H_idx]) { + rtPrecision = ov::element::f32; + } + + return rtPrecision; +} + } // namespace node } // namespace intel_cpu } // namespace ov diff --git a/src/plugins/intel_cpu/src/nodes/scaled_attn.h b/src/plugins/intel_cpu/src/nodes/scaled_attn.h index 7ce7d9f09ef..4cb09ac32d2 100644 --- a/src/plugins/intel_cpu/src/nodes/scaled_attn.h +++ b/src/plugins/intel_cpu/src/nodes/scaled_attn.h @@ -5,12 +5,14 @@ #pragma once #include #include +#include #include #include #include #include "transformations/cpu_opset/common/op/sdpa.hpp" +#include "utils/plain_tensor.hpp" namespace ov { namespace intel_cpu { @@ -41,20 +43,39 @@ public: enum KernelTypes { KT_REF, KT_ONEDNN, KT_MLAS}; + void assignState(const std::shared_ptr& state, int idx); + + const std::vector& getKVCacheOrder() const { + return m_config.config.permute_axes; + } + + ov::element::Type getKVCachePrecision(); + private: - struct Executor { - virtual void execute(dnnl::stream strm, const std::vector& inputs, const std::vector& outputs) = 0; - }; + void gatherConcatPastkv(const MemoryPtr& mem_cur_k, const MemoryPtr& mem_cur_v, const MemoryPtr& mem_beam_idx); + void updateBeamTable(const MemoryPtr& mem_beam_idx, size_t new_q_len); + void updatePastkv(const MemoryPtr& mem_cur_k, const MemoryPtr& mem_cur_v); + ov::element::Type getRuntimePrecision() const override; struct Config { ScaledDotProductAttentionWithKVCache::Config config; - bool is_concat_inplaced = false; + }; + + struct Executor { + virtual void execute(dnnl::stream strm, const Config& config, const std::vector& inputs, const MemoryPtr output, + const MemoryPtr presentk_input, const MemoryPtr presentv_input, const MemoryPtr beam_input) = 0; }; Config m_config; std::shared_ptr m_executor; template struct AttentionExecutor; - ov::element::Type rtPrecision; + friend struct ScaledDotProductAttentionKey; + + std::shared_ptr m_k_state; + std::shared_ptr m_v_state; + + ov::element::Type m_kvcache_precision = ov::element::undefined; + PlainTensor m_tmp_reorder; }; } // namespace node diff --git a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/rope_fusion.cpp b/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/rope_fusion.cpp index 0e6dbf22532..bb254160af3 100644 --- a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/rope_fusion.cpp +++ b/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/rope_fusion.cpp @@ -143,7 +143,9 @@ ov::intel_cpu::RoPEFusionCosSinPreprocess::RoPEFusionCosSinPreprocess() { auto index_Gather2 = makePattern({slice_Slice2, gather_positions_2d, 0}, {{"batch_dims", 0}}); auto unsqueeze = makePattern({index_Gather | index_Gather2, {1, 1, -1, head_dims}}); - return unsqueeze; + auto unsqueeze2 = makePattern({index_Gather2, 1}); + + return unsqueeze2 | unsqueeze; }; auto cos_tab = prepare_cos_sin_gptneox(cos_const) | prepare_cos_sin_llama(cos_const); diff --git a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_sdpa_fusion.cpp b/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_sdpa_fusion.cpp index 683609e968c..db2f39a82a9 100644 --- a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_sdpa_fusion.cpp +++ b/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_sdpa_fusion.cpp @@ -4,6 +4,8 @@ #include "stateful_sdpa_fusion.hpp" +#include + #include #include #include @@ -18,6 +20,8 @@ #include "itt.hpp" #include "ov_ops/type_relaxed.hpp" #include "transformations/cpu_opset/common/op/sdpa.hpp" +#include "utils/gen_pattern.hpp" +using namespace ov::gen_pattern; namespace ov { namespace intel_cpu { @@ -26,27 +30,80 @@ StatefulSDPAFusion::StatefulSDPAFusion() { MATCHER_SCOPE(StatefulSDPAFusion); using namespace ov::pass::pattern; - auto past_k = wrap_type(); - auto past_v = wrap_type(); - auto convert_past_k = wrap_type({past_k}); - auto convert_past_v = wrap_type({past_v}); - auto concat_input_k = std::make_shared(OutputVector{past_k, convert_past_k}); - auto concat_input_v = std::make_shared(OutputVector{past_v, convert_past_v}); - auto concat_k = wrap_type({concat_input_k, any_input()}); - auto concat_v = wrap_type({concat_input_v, any_input()}); - auto sdp0 = wrap_type({any_input(), concat_k, concat_v}); - auto sdp1 = wrap_type({any_input(), concat_k, concat_v, any_input()}); - auto sdp2 = wrap_type({any_input(), concat_k, concat_v, any_input(), any_input()}); - auto sdp = std::make_shared(OutputVector{sdp0, sdp1, sdp2}); + auto beam_idx = makePattern("i32[?]"); + auto cur_q = any_input(); + auto cur_k = any_input(); + auto cur_v = any_input(); + + auto axis_seq_len = Symbol("axis_seq_len"); + auto axis_beam = Symbol("axis_beam"); + + // past_kv can be BHLS/LBHS + auto past_k = makePattern({}); + auto past_v = makePattern({}); + + auto convert_past_k = makePattern({past_k}); + auto convert_past_v = makePattern({past_v}); + + auto gather_input_k = + makePattern({past_k | convert_past_k, beam_idx, axis_beam}, {{"batch_dims", 0}}); + auto gather_input_v = + makePattern({past_v | convert_past_v, beam_idx, axis_beam}, {{"batch_dims", 0}}); + + auto concat_k = makePattern({gather_input_k, cur_k}, {{"axis", axis_seq_len}}); + auto concat_v = makePattern({gather_input_v, cur_v}, {{"axis", axis_seq_len}}); + + auto multi_query_bcst = [](std::shared_ptr kv) { + auto reshape_kv = wrap_type({kv, any_input()}); + auto unsqueeze_kv = makePattern({kv, -2}); + auto constant_bcst = makeConst(ov::element::f32, ov::PartialShape("[...]"), [](ov::op::v0::Constant& node) { + const auto& bcst_arg = node.cast_vector(); + return std::all_of(bcst_arg.begin(), bcst_arg.end(), [](float i) { + return i == 1.0; + }); + }); + auto multiply_kv = wrap_type({reshape_kv | unsqueeze_kv, constant_bcst}); + return wrap_type({multiply_kv, any_input()}); + }; + + auto present_k = concat_k | multi_query_bcst(concat_k); + auto present_v = concat_v | multi_query_bcst(concat_v); + + // canonical q/k/v shape definition: [B,H,...L,S] + auto sdp0 = makePattern({cur_q, present_k, present_v}); + auto sdp1 = makePattern({cur_q, present_k, present_v, any_input()}); + auto sdp2 = + makePattern({cur_q, present_k, present_v, any_input(), any_input()}); + + // non-canonical q/k/v shape definitions, for example: [L, B, H, S]/[B, L, H, S] + auto order_k = wrap_type(); + auto order_v = wrap_type(); + auto order_q = wrap_type(); + auto transpose_q = makePattern({cur_q, order_q}); + auto transpose_k = makePattern({present_k, order_k}); + auto transpose_v = makePattern({present_v, order_v}); + + auto sdp_trans0 = makePattern({transpose_q, transpose_k, transpose_v}); + auto sdp_trans1 = + makePattern({transpose_q, transpose_k, transpose_v, any_input()}); + auto sdp_trans2 = makePattern( + {transpose_q, transpose_k, transpose_v, any_input(), any_input()}); + + auto sdp = sdp0 | sdp1 | sdp2 | sdp_trans0 | sdp_trans1 | sdp_trans2; ov::matcher_pass_callback callback = [=](Matcher& m) { const auto& pattern_map = m.get_pattern_value_map(); auto root = m.get_match_root(); + PatternValidator validator(m); + if (!validator) { + return false; + } + auto find_assign = [&](const ov::Output& out, opset6::Assign*& assign, opset1::Convert*& cvt) { auto present_to = out.get_target_inputs(); if (present_to.size() != 2) - return; + return false; for (auto& to : present_to) { auto to_node = to.get_node(); if (auto convert = dynamic_cast(to_node)) { @@ -58,44 +115,83 @@ StatefulSDPAFusion::StatefulSDPAFusion() { } assign = dynamic_cast(to_node); if (assign) - return; + return true; } + return false; + }; + auto check_valid_children_type = [](const ov::Output& out) { + auto children = out.get_target_inputs(); + for (auto& child : children) { + auto node = child.get_node(); + if (!one_of(node->get_type_info(), + ov::op::v13::ScaledDotProductAttention::get_type_info_static(), + ov::op::v0::ShapeOf::get_type_info_static(), + ov::op::v3::ShapeOf::get_type_info_static(), + ov::op::v0::Convert::get_type_info_static(), + ov::op::v8::Gather::get_type_info_static())) + return false; + } + return true; }; - std::shared_ptr read_cvt_k_node, read_cvt_v_node; const auto sdp_node = ov::as_type_ptr(root); const auto past_k_node = ov::as_type_ptr(pattern_map.at(past_k).get_node_shared_ptr()); const auto past_v_node = ov::as_type_ptr(pattern_map.at(past_v).get_node_shared_ptr()); + if (!check_valid_children_type(past_k_node) || !check_valid_children_type(past_v_node)) { + return false; + } const auto concat_k_node = ov::as_type_ptr(pattern_map.at(concat_k).get_node_shared_ptr()); const auto concat_v_node = ov::as_type_ptr(pattern_map.at(concat_v).get_node_shared_ptr()); - if (pattern_map.count(convert_past_k)) { - read_cvt_k_node = ov::as_type_ptr(pattern_map.at(convert_past_k).get_node_shared_ptr()); - read_cvt_v_node = ov::as_type_ptr(pattern_map.at(convert_past_v).get_node_shared_ptr()); - } - opset6::Assign* assign_k_node = nullptr, *assign_v_node = nullptr; - opset1::Convert* assign_cvt_k_node = nullptr, *assign_cvt_v_node = nullptr; - find_assign(concat_k_node, assign_k_node, assign_cvt_k_node); - if (!assign_k_node) + + opset6::Assign *assign_k_node = nullptr, *assign_v_node = nullptr; + opset1::Convert *assign_cvt_k_node = nullptr, *assign_cvt_v_node = nullptr; + if (!find_assign(concat_k_node, assign_k_node, assign_cvt_k_node)) return false; if (past_k_node->get_variable_id() != assign_k_node->get_variable_id()) return false; - find_assign(concat_v_node, assign_v_node, assign_cvt_v_node); - if (!assign_v_node) + if (!find_assign(concat_v_node, assign_v_node, assign_cvt_v_node)) return false; if (past_v_node->get_variable_id() != assign_v_node->get_variable_id()) return false; - auto args = sdp_node->input_values(); - args[1] = concat_k_node->input_value(1); - args[2] = concat_v_node->input_value(1); - args.push_back(read_cvt_k_node ? read_cvt_k_node->output(0) : past_k_node->output(0)); - args.push_back(read_cvt_v_node ? read_cvt_v_node->output(0) : past_v_node->output(0)); + // past_k & past_v must be reordered by same beam_idx + const auto gather_k_node = + ov::as_type_ptr(pattern_map.at(gather_input_k).get_node_shared_ptr()); + const auto gather_v_node = + ov::as_type_ptr(pattern_map.at(gather_input_v).get_node_shared_ptr()); + if (gather_k_node->input_value(1) != gather_v_node->input_value(1)) { + return false; + } + + OutputVector args = sdp_node->input_values(); + args[0] = pattern_map.at(cur_q); + args[1] = pattern_map.at(cur_k); + args[2] = pattern_map.at(cur_v); + args.push_back(pattern_map.at(beam_idx)); + args.push_back(gather_k_node->input_value(0)); + args.push_back(gather_v_node->input_value(0)); ov::intel_cpu::ScaledDotProductAttentionWithKVCache::Config config; config.is_causal = sdp_node->get_causal(); config.fuse_concat = true; + if (pattern_map.count(order_q) && pattern_map.count(order_k) && pattern_map.count(order_v)) { + const auto order_q_node = ov::as_type_ptr(pattern_map.at(order_q).get_node_shared_ptr()); + const auto order_k_node = ov::as_type_ptr(pattern_map.at(order_k).get_node_shared_ptr()); + const auto order_v_node = ov::as_type_ptr(pattern_map.at(order_v).get_node_shared_ptr()); + const auto& permute_q = order_q_node->cast_vector(); + const auto& permute_k = order_k_node->cast_vector(); + const auto& permute_v = order_v_node->cast_vector(); + if (permute_q != permute_k || permute_q != permute_v) { + return false; + } + config.permute_axes.resize(permute_q.size()); + for (size_t i = 0; i < permute_q.size(); i++) { + config.permute_axes[i] = static_cast(permute_q[i]); + } + } + auto old_node = sdp_node; auto new_node = std::make_shared(args, config); new_node->set_friendly_name(old_node->get_friendly_name()); diff --git a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.cpp b/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.cpp deleted file mode 100644 index e06bec71494..00000000000 --- a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.cpp +++ /dev/null @@ -1,176 +0,0 @@ -// Copyright (C) 2018-2023 Intel Corporation -// SPDX-License-Identifier: Apache-2.0 -// - -#include "stateful_transpose_sdpa_fusion.hpp" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include "itt.hpp" -#include "ov_ops/type_relaxed.hpp" -#include "transformations/cpu_opset/common/op/sdpa.hpp" - -namespace ov { -namespace intel_cpu { - -StatefulTransposeSDPAFusion::StatefulTransposeSDPAFusion() { - MATCHER_SCOPE(StatefulTransposeSDPAFusion); - using namespace ov::pass::pattern; - - auto past_k = wrap_type(); - auto past_v = wrap_type(); - auto convert_past_k = wrap_type({past_k}); - auto convert_past_v = wrap_type({past_v}); - auto concat_input_k = std::make_shared(OutputVector{past_k, convert_past_k}); - auto concat_input_v = std::make_shared(OutputVector{past_v, convert_past_v}); - auto concat_k = wrap_type({concat_input_k, any_input()}); - auto concat_v = wrap_type({concat_input_v, any_input()}); - - // multi-query branch - auto reshape_k = wrap_type({concat_k, any_input()}); - auto reshape_v = wrap_type({concat_v, any_input()}); - auto constant_k = wrap_type(); - auto constant_v = wrap_type(); - auto multiply_k = wrap_type({reshape_k, constant_k}); - auto multiply_v = wrap_type({reshape_v, constant_v}); - auto reshape1_k = wrap_type({multiply_k, any_input()}); - auto reshape1_v = wrap_type({multiply_v, any_input()}); - - auto transpose_k_input = std::make_shared(OutputVector{reshape1_k, concat_k}); - auto transpose_v_input = std::make_shared(OutputVector{reshape1_v, concat_v}); - auto order_k = wrap_type(); - auto order_v = wrap_type(); - auto transpose_k = wrap_type({transpose_k_input, order_k}); - auto transpose_v = wrap_type({transpose_v_input, order_v}); - - auto order_q = wrap_type(); - auto q_input = any_input(); - auto transpose_q = wrap_type({q_input, order_q}); - auto sdp0 = wrap_type({transpose_q, transpose_k, transpose_v}); - auto sdp1 = wrap_type({transpose_q, transpose_k, transpose_v, any_input()}); - auto sdp2 = wrap_type({transpose_q, transpose_k, transpose_v, any_input(), any_input()}); - auto sdp = std::make_shared(OutputVector{sdp0, sdp1, sdp2}); - - ov::matcher_pass_callback callback = [=](Matcher& m) { - const auto& pattern_map = m.get_pattern_value_map(); - auto root = m.get_match_root(); - auto find_assign = [&](const ov::Output& out, opset6::Assign*& assign, opset1::Convert*& cvt) { - auto present_to = out.get_target_inputs(); - if (present_to.size() != 2) - return; - for (auto& to : present_to) { - auto to_node = to.get_node(); - if (auto convert = dynamic_cast(to_node)) { - auto cvt_targets = convert->get_output_target_inputs(0); - if (cvt_targets.size() == 1) { - to_node = cvt_targets.begin()->get_node(); - cvt = convert; - } - } - assign = dynamic_cast(to_node); - if (assign) - return; - } - }; - - std::shared_ptr read_cvt_k_node, read_cvt_v_node; - const auto sdp_node = ov::as_type_ptr(root); - const auto past_k_node = ov::as_type_ptr(pattern_map.at(past_k).get_node_shared_ptr()); - const auto past_v_node = ov::as_type_ptr(pattern_map.at(past_v).get_node_shared_ptr()); - const auto concat_k_node = ov::as_type_ptr(pattern_map.at(concat_k).get_node_shared_ptr()); - const auto concat_v_node = ov::as_type_ptr(pattern_map.at(concat_v).get_node_shared_ptr()); - if (pattern_map.count(convert_past_k)) { - read_cvt_k_node = ov::as_type_ptr(pattern_map.at(convert_past_k).get_node_shared_ptr()); - read_cvt_v_node = ov::as_type_ptr(pattern_map.at(convert_past_v).get_node_shared_ptr()); - } - - // check broadcast arg has all ones - auto check_bcst = [&](const std::shared_ptr& ptr) { - const auto constant_node = ov::as_type_ptr(ptr); - const auto& bcst_arg = constant_node->cast_vector(); - return std::all_of(bcst_arg.begin(), bcst_arg.end(), [](int i) { - return i == 1.0; - }); - }; - - if (pattern_map.count(constant_k)) { - if (!check_bcst(pattern_map.at(constant_k).get_node_shared_ptr())) - return false; - } - - if (pattern_map.count(constant_v)) { - if (!check_bcst(pattern_map.at(constant_v).get_node_shared_ptr())) - return false; - } - - opset6::Assign* assign_k_node = nullptr, *assign_v_node = nullptr; - opset1::Convert* assign_cvt_k_node = nullptr, *assign_cvt_v_node = nullptr; - find_assign(concat_k_node, assign_k_node, assign_cvt_k_node); - if (!assign_k_node) - return false; - if (past_k_node->get_variable_id() != assign_k_node->get_variable_id()) - return false; - - find_assign(concat_v_node, assign_v_node, assign_cvt_v_node); - if (!assign_v_node) - return false; - if (past_v_node->get_variable_id() != assign_v_node->get_variable_id()) - return false; - auto args = sdp_node->input_values(); - args[0] = pattern_map.at(q_input).get_node_shared_ptr()->output(0); - args[1] = concat_k_node->input_value(1); - args[2] = concat_v_node->input_value(1); - args.push_back(read_cvt_k_node ? read_cvt_k_node->output(0) : past_k_node->output(0)); - args.push_back(read_cvt_v_node ? read_cvt_v_node->output(0) : past_v_node->output(0)); - ov::intel_cpu::ScaledDotProductAttentionWithKVCache::Config config; - - const auto order_q_node = ov::as_type_ptr(pattern_map.at(order_q).get_node_shared_ptr()); - const auto order_k_node = ov::as_type_ptr(pattern_map.at(order_k).get_node_shared_ptr()); - const auto order_v_node = ov::as_type_ptr(pattern_map.at(order_v).get_node_shared_ptr()); - - const auto& permute_q = order_q_node->cast_vector(); - const auto& permute_k = order_k_node->cast_vector(); - const auto& permute_v = order_v_node->cast_vector(); - if (permute_q != permute_k || permute_q != permute_v) { - return false; - } - - config.is_causal = sdp_node->get_causal(); - config.fuse_concat = true; - - config.permute_axes.resize(permute_q.size()); - for (size_t i = 0; i < permute_q.size(); i++) { - config.permute_axes[i] = static_cast(permute_q[i]); - } - auto& old_node = sdp_node; - auto new_node = std::make_shared(args, config); - new_node->set_friendly_name(old_node->get_friendly_name()); - ov::replace_node(old_node, {new_node->output(0)}); - if (assign_cvt_k_node) - assign_cvt_k_node->set_arguments({new_node->output(1)}); - else - assign_k_node->set_arguments({new_node->output(1)}); - - if (assign_cvt_v_node) - assign_cvt_v_node->set_arguments({new_node->output(2)}); - else - assign_v_node->set_arguments({new_node->output(2)}); - - return true; - }; - - auto m = std::make_shared(sdp, matcher_name); - this->register_matcher(m, callback); -} - -} // namespace intel_cpu -} // namespace ov \ No newline at end of file diff --git a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.hpp b/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.hpp deleted file mode 100644 index 94c60c36d2b..00000000000 --- a/src/plugins/intel_cpu/src/transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.hpp +++ /dev/null @@ -1,18 +0,0 @@ -// Copyright (C) 2018-2023 Intel Corporation -// SPDX-License-Identifier: Apache-2.0 -// - -#pragma once - -#include - -namespace ov { -namespace intel_cpu { -class StatefulTransposeSDPAFusion : public ov::pass::MatcherPass { -public: - OPENVINO_RTTI("StatefulTransposeSDPAFusion", "0"); - StatefulTransposeSDPAFusion(); -}; - -} // namespace intel_cpu -} // namespace ov \ No newline at end of file diff --git a/src/plugins/intel_cpu/src/transformations/transformation_pipeline.cpp b/src/plugins/intel_cpu/src/transformations/transformation_pipeline.cpp index 774b16e32af..cf961d7978c 100644 --- a/src/plugins/intel_cpu/src/transformations/transformation_pipeline.cpp +++ b/src/plugins/intel_cpu/src/transformations/transformation_pipeline.cpp @@ -114,7 +114,6 @@ #include "transformations/cpu_opset/common/pass/swap_convert_transpose.hpp" #include "transformations/cpu_opset/common/pass/rope_fusion.hpp" #include "transformations/cpu_opset/common/pass/stateful_sdpa_fusion.hpp" -#include "transformations/cpu_opset/common/pass/stateful_transpose_sdpa_fusion.hpp" // Snippets #include "snippets/pass/tokenization.hpp" @@ -662,7 +661,6 @@ void Transformations::PostLpt() { CPU_REGISTER_PASS_X64(postLPTPassManager, RoPEFusion); CPU_REGISTER_PASS_X64(postLPTPassManager, StatefulSDPAFusion); - CPU_REGISTER_PASS_X64(postLPTPassManager, StatefulTransposeSDPAFusion); postLPTPassManager.run_passes(model); } diff --git a/src/plugins/intel_cpu/src/utils/plain_tensor.hpp b/src/plugins/intel_cpu/src/utils/plain_tensor.hpp index 945e872dedf..91db670748d 100644 --- a/src/plugins/intel_cpu/src/utils/plain_tensor.hpp +++ b/src/plugins/intel_cpu/src/utils/plain_tensor.hpp @@ -118,6 +118,15 @@ struct PlainTensor { assert(i >= 0 && static_cast::type>(i) < m_rank); return m_strides[i]; } + + template + std::vector get_strides() const { + std::vector strides(m_rank); + for (size_t i = 0; i < m_rank; i++) + strides[i] = static_cast(m_strides[i]); + return strides; + } + PlainTensor(MemoryPtr mem) { reset(mem); } @@ -380,7 +389,7 @@ struct PlainTensor { } off += m_strides[i] * coordinate; } - return reinterpret_cast(m_ptr)[off]; + return (reinterpret_cast(reinterpret_cast(m_ptr) + off * m_element_size))[0]; } template diff --git a/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_multiple_query_sdp.cpp b/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_multiple_query_sdp.cpp index 3e4198aca47..1c8ad07f8fd 100644 --- a/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_multiple_query_sdp.cpp +++ b/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_multiple_query_sdp.cpp @@ -29,14 +29,15 @@ using ConcatMultiQuerySDPParams = std::tuple& obj) { - ElementType inType; + ElementType qkvType; InputShapeAndTransposeOrder inputShapeAndOrders; bool hasShapeof; - std::tie(inType, inputShapeAndOrders, hasShapeof) = obj.param; + std::tie(qkvType, inputShapeAndOrders, hasShapeof) = obj.param; std::ostringstream result; std::vector& inputShapes = inputShapeAndOrders.first; std::vector& transposeOrder = inputShapeAndOrders.second; @@ -74,8 +75,8 @@ public: } result << ")_"; } - result << "Prc=" << inType << "_"; - result << "HasShapeOf=" << hasShapeof; + result << "Prc=" << qkvType << "_"; + result << "HasShapeOf=" << hasShapeof << "_"; result << "TransposeOrder="; result << "("; for (const auto& itr : transposeOrder) { @@ -89,34 +90,34 @@ public: void SetUp() override { InputShapeAndTransposeOrder inputShapeAndOrders; bool hasShapeOf; - ElementType inType; - std::tie(inType, inputShapeAndOrders, hasShapeOf) = this->GetParam(); + ElementType qkvType; + std::tie(qkvType, inputShapeAndOrders, hasShapeOf) = this->GetParam(); std::vector& inputShapes = inputShapeAndOrders.first; std::vector& transposeOrder = inputShapeAndOrders.second; targetDevice = ov::test::utils::DEVICE_CPU; rel_threshold = 1e-2f; configuration[ov::hint::inference_precision.name()] = ov::element::f32; - if (inType == ElementType::bf16) { + if (qkvType == ElementType::bf16) { configuration[ov::hint::inference_precision.name()] = ov::element::bf16; rel_threshold = 0.01f; } init_input_shapes(inputShapes); ov::ParameterVector inputParams; // q,k,v - inputParams.push_back(std::make_shared(inType, inputDynamicShapes[0])); - inputParams.push_back(std::make_shared(inType, inputDynamicShapes[1])); - inputParams.push_back(std::make_shared(inType, inputDynamicShapes[1])); + inputParams.push_back(std::make_shared(qkvType, inputDynamicShapes[0])); + inputParams.push_back(std::make_shared(qkvType, inputDynamicShapes[1])); + inputParams.push_back(std::make_shared(qkvType, inputDynamicShapes[1])); inputParams[0]->set_friendly_name("q"); inputParams[1]->set_friendly_name("k"); inputParams[2]->set_friendly_name("v"); // pastkv init_cost - inputParams.push_back(std::make_shared(inType, inputDynamicShapes[2])); + inputParams.push_back(std::make_shared(qkvType, inputDynamicShapes[2])); auto var_k = std::make_shared( - ov::op::util::VariableInfo{inputDynamicShapes[2], inType, "pastk"}); + ov::op::util::VariableInfo{inputDynamicShapes[2], qkvType, "pastk"}); auto pastk = std::make_shared(inputParams[3], var_k); pastk->set_friendly_name("pastk_r"); auto var_v = std::make_shared( - ov::op::util::VariableInfo{inputDynamicShapes[2], inType, "pastv"}); + ov::op::util::VariableInfo{inputDynamicShapes[2], qkvType, "pastv"}); auto pastv = std::make_shared(inputParams[3], var_v); pastv->set_friendly_name("pastv_r"); std::shared_ptr pastk_shapeof, pastv_shapeof; @@ -130,14 +131,19 @@ public: auto transposeQ = std::make_shared(inputParams[0], preOrder); auto concat_axis = transposeOrder[2]; - auto concatK = std::make_shared(OutputVector{pastk, inputParams[1]}, concat_axis); - auto concatV = std::make_shared(OutputVector{pastv, inputParams[2]}, concat_axis); + auto beam_idx = std::make_shared(ElementType::i32, ov::PartialShape{-1}); + beam_idx->set_friendly_name("beam_idx"); + inputParams.push_back(beam_idx); + auto gatherK = std::make_shared(pastk, beam_idx, op::v0::Constant::create(ElementType::i32, {1}, {transposeOrder[0]})); + auto gatherV = std::make_shared(pastv, beam_idx, op::v0::Constant::create(ElementType::i32, {1}, {transposeOrder[0]})); + auto concatK = std::make_shared(OutputVector{gatherK, inputParams[1]}, concat_axis); + auto concatV = std::make_shared(OutputVector{gatherV, inputParams[2]}, concat_axis); auto unsquezeAxis = op::v0::Constant::create(ov::element::i32, {}, {-2}); auto unsqueezeK = std::make_shared(concatK, unsquezeAxis); auto unsqueezeV = std::make_shared(concatV, unsquezeAxis); - auto targetShape = op::v0::Constant::create(inType, {1, 1, 1, 4, 1}, {1}); + auto targetShape = op::v0::Constant::create(qkvType, {1, 1, 1, 4, 1}, {1}); auto broadcastK = std::make_shared(unsqueezeK, targetShape); auto broadcastV = std::make_shared(unsqueezeV, targetShape); @@ -163,13 +169,13 @@ public: const auto reshapeOrder = get_reshape_order(inputDynamicShapes[0], transposeOrder); auto postOrder = - ov::op::v0::Constant::create(ov::element::i32, {4}, std::vector{0, 2, 1, 3}); // BHLS -> BLHS + ov::op::v0::Constant::create(ov::element::i32, {4}, std::vector{2, 0, 1, 3}); // BHLS -> LBHS auto transposeSDP = std::make_shared(sdp, postOrder); auto constReshape = ov::op::v0::Constant::create(ov::element::i32, {3}, reshapeOrder); auto reshapeSDP = std::make_shared(transposeSDP, constReshape, true); // BLHS -> B,L,HxS - auto add = std::make_shared(reshapeSDP, op::v0::Constant::create(inType, {1}, {1.0f})); + auto add = std::make_shared(reshapeSDP, op::v0::Constant::create(qkvType, {1}, {1.0f})); auto pastk_assign = std::make_shared(concatK, var_k); auto pastv_assign = std::make_shared(concatV, var_v); pastk_assign->set_friendly_name("pastk_w"); @@ -208,7 +214,16 @@ public: void generate(int idx, const std::vector& targetInputStaticShapes) { inputs.clear(); auto create_input = [this](std::shared_ptr param, ov::Shape shape, float val) { - if (param->get_element_type() == element::f32) { + if (param->get_element_type() == element::i32) { + ov::Tensor t{ov::element::i32, shape}; + auto size = shape[0]; + auto* p = static_cast(t.data()); + auto start = static_cast(val); + for (size_t i = 0; i < size; i++) { + p[i] = (start + i) % size; + } + inputs.insert({param, t}); + } else if (param->get_element_type() == element::f32) { ov::Tensor t{ov::element::f32, shape}; strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); inputs.insert({param, t}); @@ -218,11 +233,12 @@ public: inputs.insert({param, t}); } }; - // q, k, v + // q, k, v, pastkv create_input(function->get_parameters()[0], targetInputStaticShapes[0], idx + 1.0f); create_input(function->get_parameters()[1], targetInputStaticShapes[1], idx + 2.0f); create_input(function->get_parameters()[2], targetInputStaticShapes[1], idx + 3.0f); create_input(function->get_parameters()[3], targetInputStaticShapes[2], idx + 4.0f); + create_input(function->get_parameters()[4], ov::Shape{targetInputStaticShapes[0][1]}, idx + 0.0f); } void prepare() { compile_model(); @@ -250,6 +266,14 @@ public: outputTensor.copy_to(copy); outputs.push_back(copy); } + auto states = inferRequest.query_state(); + for (auto&& state : states) { + auto state_tensor = state.get_state(); + ov::Tensor copy{state_tensor.get_element_type(), state_tensor.get_shape()}; + state_tensor.copy_to(copy); + outputs.push_back(copy); + } + reset(); return outputs; @@ -266,6 +290,7 @@ TEST_P(ConcatMultiQuerySDPTest, CompareWithRefs) { CheckNumberOfNodesWithType(compiledModel, "Reorder", 0); } CheckNumberOfNodesWithType(compiledModel, "Transpose", 1); + CheckNumberOfNodesWithType(compiledModel, "Gather", 0); auto expectedOutputs = run_test(functionRefs); CheckNumberOfNodesWithType(compiledModel, "ScaledDotProductAttention", 0); for (size_t i = 0; i < actualOutputs.size(); i++) { @@ -275,7 +300,7 @@ TEST_P(ConcatMultiQuerySDPTest, CompareWithRefs) { namespace { const std::vector inputShapeAndReorders = {{ - {// inputShapes ChatGLM + {// inputShapes ChatGLM, greedy search { // L1, B, H, S {{-1, 1, 8, 64}, {{10, 1, 8, 64}, {1, 1, 8, 64}, {1, 1, 8, 64}, {20, 1, 8, 64}, {1, 1, 8, 64}}}, @@ -285,7 +310,18 @@ const std::vector inputShapeAndReorders = {{ }, // transposeOrder {1, 2, 0, 3}}, + {// beam search + { + // L1, B, H, S + {{-1, -1, 8, 64}, {{10, 4, 8, 64}, {1, 4, 8, 64}, {1, 4, 8, 64}, {1, 4, 8, 64}, {1, 4, 8, 64}}}, + {{-1, -1, 2, 64}, {{10, 4, 2, 64}, {1, 4, 2, 64}, {1, 4, 2, 64}, {1, 4, 2, 64}, {1, 4, 2, 64}}}, + // L0, B, H, S + {{-1, -1, 2, 64}, {{0, 4, 2, 64}, {10, 4, 2, 64}, {11, 4, 2, 64}, {12, 4, 2, 64}, {13, 4, 2, 64}}}, + }, + // transposeOrder + {1, 2, 0, 3}}, }}; + // TODO: BF16 test is disabled due to CI machine limitation INSTANTIATE_TEST_SUITE_P(smoke_ConcatMultiQuerySDPTest, ConcatMultiQuerySDPTest, diff --git a/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_sdp.cpp b/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_sdp.cpp index 93aedea977c..f5514c62317 100644 --- a/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_sdp.cpp +++ b/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_sdp.cpp @@ -24,7 +24,8 @@ using ConcatSDPTestParams = std::tuple(pastk); pastv_shapeof = std::make_shared(pastv); } - auto concatK = std::make_shared(OutputVector{pastk, inputParams[1]}, 2); - auto concatV = std::make_shared(OutputVector{pastv, inputParams[2]}, 2); + auto beam_idx = std::make_shared(ElementType::i32, ov::PartialShape{-1}); + beam_idx->set_friendly_name("beam_idx"); + inputParams.push_back(beam_idx); + auto gatherK = std::make_shared(pastk, beam_idx, op::v0::Constant::create(ElementType::i32, {1}, {0})); + auto gatherV = std::make_shared(pastv, beam_idx, op::v0::Constant::create(ElementType::i32, {1}, {0})); + auto concatK = std::make_shared(OutputVector{gatherK, inputParams[1]}, 2); + auto concatV = std::make_shared(OutputVector{gatherV, inputParams[2]}, 2); auto sdp = std::make_shared(inputParams[0], concatK, concatV, false); sdp->set_friendly_name("mha"); auto add = std::make_shared(sdp, op::v0::Constant::create(inType, {1}, {1.0f})); @@ -141,7 +147,16 @@ public: void generate(int idx, const std::vector& targetInputStaticShapes) { inputs.clear(); auto create_input = [this] (std::shared_ptr param, ov::Shape shape, float val) { - if (param->get_element_type() == element::f32) { + if (param->get_element_type() == element::i32) { + ov::Tensor t{ov::element::i32, shape}; + auto size = shape[0]; + auto* p = static_cast(t.data()); + auto start = static_cast(val); + for (size_t i = 0; i < size; i++) { + p[i] = (start + i) % size; + } + inputs.insert({param, t}); + } else if (param->get_element_type() == element::f32) { ov::Tensor t{ov::element::f32, shape}; strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); inputs.insert({param, t}); @@ -151,11 +166,12 @@ public: inputs.insert({param, t}); } }; - // q, k, v + // q, k, v, pastkv create_input(function->get_parameters()[0], targetInputStaticShapes[0], idx + 1.0f); create_input(function->get_parameters()[1], targetInputStaticShapes[0], idx + 2.0f); create_input(function->get_parameters()[2], targetInputStaticShapes[0], idx + 3.0f); create_input(function->get_parameters()[3], targetInputStaticShapes[1], idx + 4.0f); + create_input(function->get_parameters()[4], ov::Shape{targetInputStaticShapes[0][0]}, idx + 0.0f); } void prepare() { compile_model(); @@ -194,6 +210,7 @@ TEST_P(ConcatSDPTest, CompareWithRefs) { CheckNumberOfNodesWithType(compiledModel, "ScaledDotProductAttention", 1); CheckNumberOfNodesWithType(compiledModel, "Concatenation", 0); CheckNumberOfNodesWithType(compiledModel, "Reorder", 0); + CheckNumberOfNodesWithType(compiledModel, "Gather", 0); auto expectedOutputs = run_test(functionRefs); CheckNumberOfNodesWithType(compiledModel, "ScaledDotProductAttention", 0); for (size_t i = 0; i < actualOutputs.size(); i++) { @@ -203,13 +220,20 @@ TEST_P(ConcatSDPTest, CompareWithRefs) { namespace { const std::vector> inputShapes = { - // dynamic batch + // greedy search { // B, H, L1, S {{1, 8, -1, 64}, {{1, 8, 10, 64}, {1, 8, 1, 64}, {1, 8, 1, 64}, {1, 8, 20, 64}, {1, 8, 1, 64}}}, // B, H, L0, S {{1, 8, -1, 64}, {{1, 8, 0, 64}, {1, 8, 10, 64}, {1, 8, 11, 64}, {1, 8, 12, 64}, {1, 8, 32, 64}}}, }, + // beam search + { + // B, H, L1, S + {{-1, 8, -1, 64}, {{4, 8, 10, 64}, {4, 8, 1, 64}, {4, 8, 1, 64}, {4, 8, 1, 64}, {4, 8, 1, 64}}}, + // B, H, L0, S + {{-1, 8, -1, 64}, {{4, 8, 0, 64}, {4, 8, 10, 64}, {4, 8, 11, 64}, {4, 8, 12, 64}, {4, 8, 13, 64}}}, + }, }; INSTANTIATE_TEST_SUITE_P(smoke_ConcatSDPTest, diff --git a/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_transpose_sdp_transpose.cpp b/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_transpose_sdp_transpose.cpp index ea6d37e9e65..2eddaa63050 100644 --- a/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_transpose_sdp_transpose.cpp +++ b/src/plugins/intel_cpu/tests/functional/subgraph_tests/src/concat_transpose_sdp_transpose.cpp @@ -29,7 +29,8 @@ using ConcatSDPTransposeTestParams = std::tuple, - virtual public ov::test::SubgraphBaseTest, - public CPUTestsBase { +class ConcatSDPTransposeTestBase : public testing::WithParamInterface, + virtual public ov::test::SubgraphBaseTest, + public CPUTestsBase { public: static std::string getTestCaseName(const testing::TestParamInfo& obj) { ElementType inType; @@ -90,7 +91,7 @@ public: bool hasShapeOf; std::tie(inType, inputShapeAndOrders, hasShapeOf) = this->GetParam(); std::vector& inputShapes = inputShapeAndOrders.first; - std::vector& transposeOrder = inputShapeAndOrders.second; + transposeOrder = inputShapeAndOrders.second; targetDevice = ov::test::utils::DEVICE_CPU; rel_threshold = 1e-2f; configuration[ov::hint::inference_precision.name()] = ov::element::f32; @@ -128,8 +129,13 @@ public: auto transposeQ = std::make_shared(inputParams[0], preOrder); auto concat_axis = transposeOrder[2]; - auto concatK = std::make_shared(OutputVector{pastk, inputParams[1]}, concat_axis); - auto concatV = std::make_shared(OutputVector{pastv, inputParams[2]}, concat_axis); + auto beam_idx = std::make_shared(ElementType::i32, ov::PartialShape{-1}); + beam_idx->set_friendly_name("beam_idx"); + inputParams.push_back(beam_idx); + auto gatherK = std::make_shared(pastk, beam_idx, op::v0::Constant::create(ElementType::i32, {1}, {0})); + auto gatherV = std::make_shared(pastv, beam_idx, op::v0::Constant::create(ElementType::i32, {1}, {0})); + auto concatK = std::make_shared(OutputVector{gatherK, inputParams[1]}, concat_axis); + auto concatV = std::make_shared(OutputVector{gatherV, inputParams[2]}, concat_axis); auto transposeK = std::make_shared(concatK, preOrder); auto transposeV = std::make_shared(concatV, preOrder); @@ -183,7 +189,7 @@ public: SubgraphBaseTest::generate_inputs(shapes); } template - void strided_iota(IT first, size_t n, T value, T stride) { + static void strided_iota(IT first, size_t n, T value, T stride) { for (size_t i = 0; i < n; i++) { *first++ = value; value += stride; @@ -191,22 +197,33 @@ public: } void generate(int idx, const std::vector& targetInputStaticShapes) { inputs.clear(); - auto create_input = [this](std::shared_ptr param, ov::Shape shape, float val) { - if (param->get_element_type() == element::f32) { + auto create_input = [this] (std::shared_ptr param, ov::Shape shape, float val) { + if (param->get_element_type() == element::i32) { + ov::Tensor t{ov::element::i32, shape}; + auto size = shape[0]; + auto* p = static_cast(t.data()); + auto start = static_cast(val); + for (size_t i = 0; i < size; i++) { + p[i] = (start + i) % size; + } + inputs.insert({param, t}); + } else if (param->get_element_type() == element::f32) { ov::Tensor t{ov::element::f32, shape}; strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); inputs.insert({param, t}); } else { + ASSERT_TRUE(param->get_element_type() == element::bf16); ov::Tensor t{ov::element::bf16, shape}; strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); inputs.insert({param, t}); } }; - // q, k, v + // q, k, v, pastkv create_input(function->get_parameters()[0], targetInputStaticShapes[0], idx + 1.0f); create_input(function->get_parameters()[1], targetInputStaticShapes[0], idx + 2.0f); create_input(function->get_parameters()[2], targetInputStaticShapes[0], idx + 3.0f); create_input(function->get_parameters()[3], targetInputStaticShapes[1], idx + 4.0f); + create_input(function->get_parameters()[4], ov::Shape{targetInputStaticShapes[0][0]}, idx + 0.0f); } void prepare() { compile_model(); @@ -218,6 +235,11 @@ public: state.reset(); } } + std::vector transposeOrder; +}; + +class ConcatSDPTransposeTest : public ConcatSDPTransposeTestBase { +public: std::vector run_test(std::shared_ptr model) { function = model; prepare(); @@ -234,6 +256,14 @@ public: outputTensor.copy_to(copy); outputs.push_back(copy); } + auto states = inferRequest.query_state(); + for (auto&& state : states) { + auto state_tensor = state.get_state(); + ov::Tensor copy{state_tensor.get_element_type(), state_tensor.get_shape()}; + state_tensor.copy_to(copy); + outputs.push_back(copy); + } + reset(); return outputs; @@ -246,6 +276,7 @@ TEST_P(ConcatSDPTransposeTest, CompareWithRefs) { CheckNumberOfNodesWithType(compiledModel, "Concatenation", 0); CheckNumberOfNodesWithType(compiledModel, "Reorder", 0); CheckNumberOfNodesWithType(compiledModel, "Transpose", 1); + CheckNumberOfNodesWithType(compiledModel, "Gather", 0); auto expectedOutputs = run_test(functionRefs); CheckNumberOfNodesWithType(compiledModel, "ScaledDotProductAttention", 0); for (size_t i = 0; i < actualOutputs.size(); i++) { @@ -256,33 +287,27 @@ TEST_P(ConcatSDPTransposeTest, CompareWithRefs) { namespace { const std::vector inputShapeAndReorders = { { - // inputShapes LLama - { - // B, H, L1, S - {{1, 8, -1, 64}, {{1, 8, 10, 64}, {1, 8, 1, 64}, {1, 8, 1, 64}, {1, 8, 20, 64}, {1, 8, 1, 64}}}, - // B, H, L0, S - {{1, 8, -1, 64}, {{1, 8, 0, 64}, {1, 8, 10, 64}, {1, 8, 11, 64}, {1, 8, 12, 64}, {1, 8, 32, 64}}}, - }, - // transposeOrder - {0, 1, 2, 3}}, - {// inputShapes QWen - { - // B, L1, H, S - {{1, -1, 8, 64}, {{1, 10, 8, 64}, {1, 1, 8, 64}, {1, 1, 8, 64}, {1, 20, 8, 64}, {1, 1, 8, 64}}}, - // B, L0, H, S - {{1, -1, 8, 64}, {{1, 0, 8, 64}, {1, 10, 8, 64}, {1, 11, 8, 64}, {1, 12, 8, 64}, {1, 32, 8, 64}}}, - }, - // transposeOrder - {0, 2, 1, 3}}, - {// inputShapes ChatGLM - { - // L1, B, H, S - {{-1, 1, 8, 64}, {{10, 1, 8, 64}, {1, 1, 8, 64}, {1, 1, 8, 64}, {20, 1, 8, 64}, {1, 1, 8, 64}}}, - // L0, B, H, S - {{-1, 1, 8, 64}, {{0, 1, 8, 64}, {10, 1, 8, 64}, {11, 1, 8, 64}, {12, 1, 8, 64}, {32, 1, 8, 64}}}, - }, - // transposeOrder - {1, 2, 0, 3}}, + // greedy search + {{ + // B, L1, H, S + {{1, -1, 8, 64}, {{1, 10, 8, 64}, {1, 1, 8, 64}, {1, 1, 8, 64}, {1, 20, 8, 64}, {1, 1, 8, 64}}}, + // B, L0, H, S + {{1, -1, 8, 64}, {{1, 0, 8, 64}, {1, 10, 8, 64}, {1, 11, 8, 64}, {1, 12, 8, 64}, {1, 32, 8, 64}}}, + }, + // transposeOrder + {0, 2, 1, 3} + }, + // beam search + {{ + // B, L1, H, S + {{-1, -1, 8, 64}, {{4, 10, 8, 64}, {4, 1, 8, 64}, {4, 1, 8, 64}, {4, 1, 8, 64}, {4, 1, 8, 64}}}, + // B, L0, H, S + {{-1, -1, 8, 64}, {{4, 0, 8, 64}, {4, 10, 8, 64}, {4, 11, 8, 64}, {4, 12, 8, 64}, {4, 13, 8, 64}}}, + }, + // transposeOrder + {0, 2, 1, 3} + } + } }; INSTANTIATE_TEST_SUITE_P(smoke_ConcatSDPTransposeTest, @@ -291,6 +316,125 @@ INSTANTIATE_TEST_SUITE_P(smoke_ConcatSDPTransposeTest, ::testing::ValuesIn(inputShapeAndReorders), ::testing::Values(true, false)), ConcatSDPTransposeTest::getTestCaseName); +} // namespace + +class ConcatSDPTransposeTestSetState : public ConcatSDPTransposeTestBase { +public: + void reduce_state() { + auto states = inferRequest.query_state(); + for (auto&& state : states) { + auto state_tensor = state.get_state(); + ov::Tensor copy{state_tensor.get_element_type(), state_tensor.get_shape()}; + state_tensor.copy_to(copy); + auto new_shape = state_tensor.get_shape(); + ASSERT_GE(new_shape[transposeOrder[2]], 1); + new_shape[transposeOrder[2]] -= 1; + ov::Tensor new_state{state_tensor.get_element_type(), new_shape, copy.data()}; + state.set_state(new_state); + } + } + void new_state(ov::element::Type& type, const ov::Shape& pastKVInitShape) { + auto fill = [] (ov::Tensor& t, float val) { + auto shape = t.get_shape(); + if (t.get_element_type() == element::f32) { + strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); + } else if (t.get_element_type() == element::f16) { + strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); + } else { + ASSERT_TRUE(t.get_element_type() == element::bf16); + strided_iota(static_cast(t.data()), t.get_size(), val, 0.1f); + } + }; + float val = 0; + auto states = inferRequest.query_state(); + for (auto&& state : states) { + auto new_shape = pastKVInitShape; + new_shape[transposeOrder[2]] = 3; + ov::Tensor new_state{type, new_shape}; + fill(new_state, val); + val += 0.13f; + + state.set_state(new_state); + } + } + std::vector run_test(std::shared_ptr model) { + function = model; + prepare(); + std::vector outputs; + // case 1: initialization + pastkv reaches limitation, remove some state + int idx = 0; + for (auto&& shapes : targetStaticShapes) { + generate(idx++, shapes); + for (const auto& input : inputs) { + inferRequest.set_tensor(input.first, input.second); + } + inferRequest.infer(); + auto outputTensor = inferRequest.get_output_tensor(0); + ov::Tensor copy{outputTensor.get_element_type(), outputTensor.get_shape()}; + outputTensor.copy_to(copy); + outputs.push_back(copy); + if (idx > 1) { + reduce_state(); + } + } + + // case 2: after reset, set_state at once + auto pastKVType = inferRequest.query_state()[0].get_state().get_element_type(); + reset(); + new_state(pastKVType, targetStaticShapes[0][1]); + idx = 0; + for (auto&& shapes : targetStaticShapes) { + generate(idx++, shapes); + for (const auto& input : inputs) { + inferRequest.set_tensor(input.first, input.second); + } + inferRequest.infer(); + auto outputTensor = inferRequest.get_output_tensor(0); + ov::Tensor copy{outputTensor.get_element_type(), outputTensor.get_shape()}; + outputTensor.copy_to(copy); + outputs.push_back(copy); + } + + return outputs; + } +}; + +TEST_P(ConcatSDPTransposeTestSetState, CompareWithRefs) { + auto actualOutputs = run_test(function); + CheckNumberOfNodesWithType(compiledModel, "ScaledDotProductAttention", 1); + CheckNumberOfNodesWithType(compiledModel, "Concatenation", 0); + CheckNumberOfNodesWithType(compiledModel, "Reorder", 0); + CheckNumberOfNodesWithType(compiledModel, "Transpose", 1); + CheckNumberOfNodesWithType(compiledModel, "Gather", 0); + auto expectedOutputs = run_test(functionRefs); + CheckNumberOfNodesWithType(compiledModel, "ScaledDotProductAttention", 0); + for (size_t i = 0; i < actualOutputs.size(); i++) { + ov::test::utils::compare(expectedOutputs[i], actualOutputs[i], abs_threshold, rel_threshold); + } +} + +namespace { +const std::vector inputShapeAndReordersSetState = { + { + // beam search + {{ + // B, L1, H, S + {{-1, -1, 8, 64}, {{4, 10, 8, 64}, {4, 1, 8, 64}, {4, 1, 8, 64}, {4, 1, 8, 64}}}, + // B, L0, H, S and init tensor + {{-1, -1, 8, 64}, {{4, 2, 8, 64}, {4, 12, 8, 64}, {4, 13, 8, 64}, {4, 14, 8, 64}}}, + }, + // transposeOrder + {0, 2, 1, 3} + } + } +}; + +INSTANTIATE_TEST_SUITE_P(smoke_ConcatSDPTransposeTestSetState, + ConcatSDPTransposeTestSetState, + ::testing::Combine(::testing::Values(ElementType::f32), + ::testing::ValuesIn(inputShapeAndReordersSetState), + ::testing::Values(false)), + ConcatSDPTransposeTest::getTestCaseName); } // namespace } // namespace SubgraphTestsDefinitions diff --git a/src/plugins/intel_cpu/tests/unit/graph/scaled_attn.cpp b/src/plugins/intel_cpu/tests/unit/graph/scaled_attn.cpp deleted file mode 100644 index 48db2d5fd70..00000000000 --- a/src/plugins/intel_cpu/tests/unit/graph/scaled_attn.cpp +++ /dev/null @@ -1,177 +0,0 @@ -// Copyright (C) 2018-2023 Intel Corporation -// SPDX-License-Identifier: Apache-2.0 -// -#include - -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#include -#include -#include - -#include "common_test_utils/common_utils.hpp" -#include "cache/multi_cache.h" -#include "ov_models/builders.hpp" -#include "nodes/scaled_attn.h" -#include "nodes/input.h" -#include "nodes/convert.h" -#include "graph.h" -#include "cpu_tensor.h" - -using namespace ov::intel_cpu; - -#ifdef OPENVINO_ARCH_ARM64 -// Ticket: 126975 -TEST(ScaledAttnGraphTest, DISABLED_smoke_Check_Scaled_Concat_Noplace) { -#else -TEST(ScaledAttnGraphTest, smoke_Check_Scaled_Concat_Noplace) { -#endif - auto build_graph = [](const ov::Shape& shape, float* qkv_val, float* past_kv_val) { - auto qkv = ov::op::v0::Constant::create(ov::element::f32, shape, qkv_val); - qkv->set_friendly_name("qkv_const"); - auto pastkv_f32 = ov::op::v0::Constant::create(ov::element::f32, shape, past_kv_val); - pastkv_f32->set_friendly_name("pastkv_const_f32"); - auto pastkv = std::make_shared(pastkv_f32, ov::element::f16); - pastkv->set_friendly_name("pastkv_const"); - // only need a dynamic parameter but its value will not be used - auto attn = std::make_shared(ov::element::f32, ov::PartialShape{-1}); - attn->set_friendly_name("attn"); - - ov::intel_cpu::ScaledDotProductAttentionWithKVCache::Config config; - config.fuse_concat = true; - config.is_causal = true; - auto sdpa = std::make_shared(ov::OutputVector{qkv, qkv, qkv, attn, pastkv, pastkv}, config); - auto out_pastk_convert = std::make_shared(sdpa->output(1), ov::element::f32); - auto out_pastv_convert = std::make_shared(sdpa->output(2), ov::element::f32); - auto out_qkv = std::make_shared(sdpa->output(0)); - out_qkv->set_friendly_name("qkv"); - auto out_pastk = std::make_shared(out_pastk_convert); - out_pastk->set_friendly_name("pastk"); - auto out_pastv = std::make_shared(out_pastv_convert); - out_pastv->set_friendly_name("pastv"); - - std::unordered_set nodes_set; - std::vector graph_edges; - - auto add_edge = [&](const NodePtr& parent, const NodePtr& child, size_t parentPort, size_t childPort) -> void { - auto edge = std::make_shared(parent, child, parentPort, childPort); - child->addEdge(edge); - graph_edges.push_back(edge); - nodes_set.insert(parent); - nodes_set.insert(child); - }; - - //create graph context - Config conf; - conf.rtCacheCapacity = 0; - auto context = std::make_shared(conf, nullptr, nullptr, false); - - auto qkv_node = std::make_shared(qkv, context); - auto pastkv_f32_node = std::make_shared(pastkv_f32, context); - auto attn_node = std::make_shared(attn, context); - auto pastkv_node = std::make_shared(pastkv, context); - auto sdpa_node = std::make_shared(sdpa, context); - auto out_pastk_node_convert = std::make_shared(out_pastk_convert, context); - auto out_pastv_node_convert = std::make_shared(out_pastv_convert, context); - auto out_qkv_node = std::make_shared(out_qkv, context); - auto out_pastk_node = std::make_shared(out_pastk, context); - auto out_pastv_node = std::make_shared(out_pastv, context); - - add_edge(qkv_node, sdpa_node, 0, 0); - add_edge(qkv_node, sdpa_node, 0, 1); - add_edge(qkv_node, sdpa_node, 0, 2); - add_edge(attn_node, sdpa_node, 0, 3); - add_edge(pastkv_f32_node, pastkv_node, 0, 0); - add_edge(pastkv_node, sdpa_node, 0, 4); - add_edge(pastkv_node, sdpa_node, 0, 5); - add_edge(sdpa_node, out_qkv_node, 0, 0); - add_edge(sdpa_node, out_pastk_node_convert, 1, 0); - add_edge(sdpa_node, out_pastv_node_convert, 2, 0); - add_edge(out_pastk_node_convert, out_pastk_node, 0, 0); - add_edge(out_pastv_node_convert, out_pastv_node, 0, 0); - - std::vector graph_nodes(nodes_set.begin(), nodes_set.end()); - - Graph graph; - graph.CreateGraph(graph_nodes, graph_edges, context, "test_graph"); - return graph; - }; - - auto run_graph = [] (Graph& graph) { - graph.GetInputNodesMap().begin()->second->redefineOutputMemory(0, {1}); - - for (auto& node : graph.GetNodes()) { - if (node->isDynamicNode()) { - node->updateShapes(); - node->updateDynamicParams(); - } - } - graph.Infer(); - }; - - auto check_graph = [] (Graph& graph, std::map>& expected) { - auto& outputNodesMap = graph.GetOutputNodesMap(); - auto is_same = [] (float a, float b) { - return std::abs(a - b) < 0.01f; - }; - for (auto &outputMap : outputNodesMap) { - auto name = outputMap.first; - if (expected.count(name) == 0) { - continue; - } - auto node = outputMap.second; - auto parentEdge = node->getParentEdgeAt(0); - const auto& memory = parentEdge->getMemoryPtr(); - auto size = memory->getSize() / sizeof(float); - auto p = reinterpret_cast(memory->getData()); - for (size_t i = 0; i < size; i++) { - ASSERT_EQ(is_same(p[i], expected.at(name).first[i]), true); - } - ASSERT_EQ(memory->getShape(), ov::intel_cpu::Shape(expected.at(name).second)); - } - }; - - auto find_node_type = [](const Graph& graph, Type type) -> NodePtr { - auto&& nodes = graph.GetNodes(); - auto itr = - std::find_if(nodes.begin(), nodes.end(), [=](const NodePtr& node){ return type == node->getType(); }); - - if (itr == nodes.end()) { - return nullptr; - } - - return (*itr); - }; - - auto strided_iota = [] (float* first, size_t n, float value, float stride) { - for (size_t i = 0; i < n; i++) { - *first++ = value; - value += stride; - } - }; - - ov::Shape shape{1, 1, 8, 8}; - const size_t elements_count = std::accumulate(shape.begin(), shape.end(), size_t{1}, std::multiplies()); - std::vector val(elements_count * 2); - strided_iota(val.data(), val.size(), -10.0f, 0.1f); - auto graph = build_graph(shape, val.data() + elements_count, val.data()); - run_graph(graph); - // if no inplace, the pastk and pastv will concat, check shape and value - ov::Shape expectedShape(shape); - expectedShape[2] *= 2; - std::map> expected{ - {"pastk", std::make_pair(val.data(), expectedShape)}, - {"pastv", std::make_pair(val.data(), expectedShape)}}; - check_graph(graph, expected); - auto spd = find_node_type(graph, Type::ScaledDotProductAttention)->getSelectedPrimitiveDescriptor(); - ASSERT_EQ(spd->getConfig().outConfs[1].inPlace(), -1); - ASSERT_EQ(spd->getConfig().outConfs[2].inPlace(), -1); -} \ No newline at end of file diff --git a/src/plugins/intel_cpu/tests/unit/transformations/state_concat_sdpa.cpp b/src/plugins/intel_cpu/tests/unit/transformations/state_concat_sdpa.cpp index 1ce6263cb2d..ffa562ba4d8 100644 --- a/src/plugins/intel_cpu/tests/unit/transformations/state_concat_sdpa.cpp +++ b/src/plugins/intel_cpu/tests/unit/transformations/state_concat_sdpa.cpp @@ -26,6 +26,7 @@ static std::shared_ptr makeSDPA(const ov::PartialShape& inputShape, b auto k = std::make_shared(element::f32, inputShape); auto v = std::make_shared(element::f32, inputShape); auto init = std::make_shared(element::f32, inputShape); + auto beam_idx = std::make_shared(element::i32, ov::PartialShape{-1}); auto var_k = std::make_shared( ov::op::util::VariableInfo{inputShape, element::f32, "pastk"}); std::shared_ptr pastk = std::make_shared(k, var_k); @@ -40,11 +41,13 @@ static std::shared_ptr makeSDPA(const ov::PartialShape& inputShape, b if (isRef) { ov::intel_cpu::ScaledDotProductAttentionWithKVCache::Config config; config.fuse_concat = true; - auto new_node = std::make_shared(OutputVector{q, k, v, pastk, pastv}, config); + auto new_node = std::make_shared(OutputVector{q, k, v, beam_idx, pastk, pastv}, config); sdp = new_node->output(0); concatK = new_node->output(1); concatV = new_node->output(2); } else { + pastk = std::make_shared(pastk, beam_idx, op::v0::Constant::create(element::i32, {1}, {0})); + pastv = std::make_shared(pastv, beam_idx, op::v0::Constant::create(element::i32, {1}, {0})); concatK = std::make_shared(OutputVector{pastk, k}, 2); concatV = std::make_shared(OutputVector{pastv, v}, 2); sdp = std::make_shared(q, concatK, concatV, false); @@ -59,7 +62,7 @@ static std::shared_ptr makeSDPA(const ov::PartialShape& inputShape, b ResultVector results{std::make_shared(add)}; SinkVector sinks{pastk_assign, pastv_assign}; - return std::make_shared(results, sinks, ParameterVector{q, k, v, init}, "ConcatSDP"); + return std::make_shared(results, sinks, ParameterVector{q, k, v, init, beam_idx}, "ConcatSDP"); } TEST(TransformationTests, StateConcatSDPA) {