From 2d0e0366f7b3f49d9400c56581ec738c491fe923 Mon Sep 17 00:00:00 2001 From: Alexandra Sidorova Date: Thu, 30 May 2024 17:07:46 +0300 Subject: [PATCH] [Snippets] United Static and Dynamic Loops into one node (#24525) ### Details: - *United `LoopEndStatic` and `LoopEndDynamic` into one node `LoopEnd` to avoid extra conditions in the code and improve performance since some pointer data shifts might be known and compiled in JIT code* - *Removed dynamic aarch64 loop emitters since they don't work anyway CVS-141550* - *Added support dynamism to `IdentifyBuffers` and `DefineBufferClusters`. It's not efficient algorithm since we don't know exact values of data pointer shifts and cannot be sure that they will be proportionally in runtime. It should be implemented as the separate feature based on some judgments, for example* ### Tickets: - *141268* ### Prerequisites: - *https://github.com/openvinotoolkit/openvino/pull/21922* --- .../include/snippets/lowered/loop_info.hpp | 20 ++ .../lowered/pass/identify_buffers.hpp | 6 + .../snippets/lowered/pass/insert_loops.hpp | 1 - .../snippets/include/snippets/op/loop.hpp | 71 ++--- src/common/snippets/src/lowered/loop_info.cpp | 21 +- .../pass/clean_repeated_ptr_shifts.cpp | 22 +- .../src/lowered/pass/cleanup_loop_offsets.cpp | 9 +- .../lowered/pass/define_buffer_clusters.cpp | 46 ++-- .../src/lowered/pass/identify_buffers.cpp | 19 +- .../snippets/src/lowered/pass/init_loops.cpp | 13 +- .../src/lowered/pass/insert_loops.cpp | 43 +-- .../pass/insert_specific_iterations.cpp | 8 +- .../src/lowered/pass/iter_handler.cpp | 6 +- .../pass/optimize_loop_single_evaluation.cpp | 7 +- .../snippets/src/lowered/pass/split_loops.cpp | 4 +- .../snippets/src/lowered/pass/validate.cpp | 46 +--- .../lowered/pass/validate_expanded_loops.cpp | 14 +- src/common/snippets/src/op/loop.cpp | 227 ++++++++-------- .../src/shape_inference/shape_inference.cpp | 6 +- .../snippets/tests/src/lowered/pass/loop.cpp | 3 +- .../snippets/tests/src/lowering_utils.cpp | 6 +- .../snippets/aarch64/cpu_generator.cpp | 6 +- .../snippets/aarch64/jit_loop_emitters.cpp | 179 ++----------- .../snippets/aarch64/jit_loop_emitters.hpp | 84 +----- .../emitters/snippets/x64/cpu_generator.cpp | 6 +- .../snippets/x64/jit_loop_emitters.cpp | 250 +++++++----------- .../snippets/x64/jit_loop_emitters.hpp | 95 ++----- src/plugins/intel_cpu/src/extension.cpp | 6 +- 28 files changed, 435 insertions(+), 789 deletions(-) diff --git a/src/common/snippets/include/snippets/lowered/loop_info.hpp b/src/common/snippets/include/snippets/lowered/loop_info.hpp index ca28b27a760..2bd1a4f8bab 100644 --- a/src/common/snippets/include/snippets/lowered/loop_info.hpp +++ b/src/common/snippets/include/snippets/lowered/loop_info.hpp @@ -34,6 +34,12 @@ public: */ virtual std::shared_ptr clone_with_new_expr(const ExpressionMap& expr_map) const = 0; + /** + * @brief Check if some parameters of Loop are dynamic (undefined) + * @return True if some parameters of Loop are unknown, False if all parameters are static + */ + virtual bool is_dynamic() const; + /** * @brief Returns count of input ports * @return count @@ -184,6 +190,8 @@ public: int64_t ptr_increment = 0; int64_t finalization_offset = 0; int64_t data_size = 0; + + bool is_dynamic() const; }; // The structure describes full information about port // - TODO [140365] : UnifiedLoopInfo should have the map of LoopPorts and LoopDesc as class field @@ -212,6 +220,12 @@ public: */ std::shared_ptr clone_with_new_expr(const ExpressionMap& expr_map) const override; + /** + * @brief Check if some parameters of Loop are dynamic (undefined) + * @return True if some parameters of Loop are unknown, False if all parameters are static + */ + bool is_dynamic() const override; + /** * @brief Returns handlers of loop specific iterations * @return m_handlers @@ -373,6 +387,12 @@ public: */ std::shared_ptr clone_with_new_expr(const ExpressionMap& expr_map) const override; + /** + * @brief Check if some parameters of Loop are dynamic (undefined) + * @return True if some parameters of Loop are unknown, False if all parameters are static + */ + bool is_dynamic() const override; + /** * @brief Returns original unified LoopInfo from which this LoopInfo was created * @return const reference of m_unified_loop_info diff --git a/src/common/snippets/include/snippets/lowered/pass/identify_buffers.hpp b/src/common/snippets/include/snippets/lowered/pass/identify_buffers.hpp index 31631b9b0ec..2289ef0246e 100644 --- a/src/common/snippets/include/snippets/lowered/pass/identify_buffers.hpp +++ b/src/common/snippets/include/snippets/lowered/pass/identify_buffers.hpp @@ -6,6 +6,8 @@ #include "pass.hpp" +#include "snippets/utils.hpp" + namespace ov { namespace snippets { namespace lowered { @@ -46,6 +48,10 @@ public: int64_t ptr_increment = 0; int64_t finalization_offset = 0; + inline bool is_static() const { + return !utils::is_dynamic_value(ptr_increment) && !utils::is_dynamic_value(finalization_offset); + } + friend bool operator==(const ShiftPtrParams& lhs, const ShiftPtrParams& rhs); friend bool operator!=(const ShiftPtrParams& lhs, const ShiftPtrParams& rhs); }; diff --git a/src/common/snippets/include/snippets/lowered/pass/insert_loops.hpp b/src/common/snippets/include/snippets/lowered/pass/insert_loops.hpp index 1329fa22a6b..1c86ccbbc83 100644 --- a/src/common/snippets/include/snippets/lowered/pass/insert_loops.hpp +++ b/src/common/snippets/include/snippets/lowered/pass/insert_loops.hpp @@ -26,7 +26,6 @@ public: bool run(LinearIR& linear_ir, lowered::LinearIR::constExprIt begin, lowered::LinearIR::constExprIt end) override; private: static void insertion(LinearIR& linear_ir, const LoopManagerPtr& loop_manager, size_t loop_id); - static bool is_loop_dynamic(const UnifiedLoopInfoPtr& loop_info); }; } // namespace pass diff --git a/src/common/snippets/include/snippets/op/loop.hpp b/src/common/snippets/include/snippets/op/loop.hpp index 1053dba7d2a..2226110555b 100644 --- a/src/common/snippets/include/snippets/op/loop.hpp +++ b/src/common/snippets/include/snippets/op/loop.hpp @@ -40,26 +40,13 @@ public: LoopBegin(); void validate_and_infer_types() override; + std::shared_ptr clone_with_new_inputs(const OutputVector& inputs) const override; std::shared_ptr get_loop_end() const; protected: void validate_and_infer_types_except_LoopEnd(); }; -class LoopBeginStatic : public LoopBegin { -public: - OPENVINO_OP("LoopBeginStatic", "SnippetsOpset", LoopBegin); - LoopBeginStatic() = default; - std::shared_ptr clone_with_new_inputs(const OutputVector& inputs) const override; -}; - -class LoopBeginDynamic : public LoopBegin { -public: - OPENVINO_OP("LoopBeginDynamic", "SnippetsOpset", LoopBegin); - LoopBeginDynamic() = default; - std::shared_ptr clone_with_new_inputs(const OutputVector& inputs) const override; -}; - /** * @interface LoopEnd * @brief Marks the end of the Loop region and defines the loop properties. @@ -77,78 +64,50 @@ class LoopEnd : public LoopBase { public: OPENVINO_OP("LoopEnd", "SnippetsOpset", LoopBase); LoopEnd() = default; - LoopEnd(const Output& loop_begin, size_t work_amount_increment, std::vector is_incremented, + LoopEnd(const Output& loop_begin, size_t work_amount, size_t work_amount_increment, + std::vector is_incremented, std::vector ptr_increments, std::vector finalization_offsets, std::vector element_type_sizes, size_t input_num, size_t output_num, size_t id); void validate_and_infer_types() override; bool visit_attributes(AttributeVisitor& visitor) override; + std::shared_ptr clone_with_new_inputs(const OutputVector& inputs) const override; + std::shared_ptr get_loop_begin(); const std::vector& get_is_incremented() const; + const std::vector& get_finalization_offsets() const; + const std::vector& get_ptr_increments() const; const std::vector& get_element_type_sizes() const; + size_t get_work_amount() const; size_t get_increment() const; size_t get_id() const; size_t get_input_num() const; size_t get_output_num() const; bool get_evaluate_once() const; + bool has_dynamic_params() const; void set_is_incremented(std::vector is_incremented); + void set_finalization_offsets(std::vector offsets); + void set_ptr_increments(std::vector new_ptr_increments); + void set_work_amount(size_t new_work_amount); void set_increment(size_t new_increment); void set_evaluate_once(bool once); void set_id(size_t id); protected: std::vector m_is_incremented = {}; + std::vector m_ptr_increments = {}; + std::vector m_finalization_offsets = {}; std::vector m_element_type_sizes = {}; + size_t m_work_amount = 0; size_t m_work_amount_increment = 0; size_t m_input_num = 0; size_t m_output_num = 0; size_t m_id = 0; // the corresponding Loop identificator in LoopManager -}; -class LoopEndStatic : public LoopEnd { -public: - OPENVINO_OP("LoopEndStatic", "SnippetsOpset", LoopEnd); - LoopEndStatic() = default; - LoopEndStatic(const Output& loop_begin, size_t work_amount, size_t work_amount_increment, - std::vector is_incremented, std::vector ptr_increments, std::vector finalization_offsets, - std::vector element_type_sizes, size_t input_num, size_t output_num, size_t id); - std::shared_ptr clone_with_new_inputs(const OutputVector& inputs) const override; - - void validate_and_infer_types() override; - bool visit_attributes(AttributeVisitor& visitor) override; - - // update_ptr_increments resets non-zero increments to the new_increments. It's used when work_amount_increment is - // updated and we need to refresh ptr increments accordingly while respecting the broadcasting pattern - void update_ptr_increments(int64_t new_increment); - - const std::vector& get_finalization_offsets() const; - const std::vector& get_ptr_increments() const; - size_t get_work_amount() const; - bool get_evaluate_once() const; - - void set_finalization_offsets(std::vector offsets); - void set_ptr_increments(std::vector new_ptr_increments); - void set_work_amount(size_t new_work_amount); - void set_evaluate_once(bool once); - -protected: - std::vector m_ptr_increments = {}; - std::vector m_finalization_offsets = {}; - size_t m_work_amount = 0; bool m_evaluate_once = false; // true if the Loop is executed only once, used to skip setting and testing the loop counter }; -class LoopEndDynamic : public LoopEnd { -public: - OPENVINO_OP("LoopEndDynamic", "SnippetsOpset", LoopEnd); - LoopEndDynamic() = default; - LoopEndDynamic(const Output& loop_begin, size_t work_amount_increment, std::vector is_incremented, - std::vector element_type_sizes, size_t input_num, size_t output_num, size_t id); - - std::shared_ptr clone_with_new_inputs(const OutputVector& inputs) const override; -}; - } // namespace op } // namespace snippets } // namespace ov diff --git a/src/common/snippets/src/lowered/loop_info.cpp b/src/common/snippets/src/lowered/loop_info.cpp index e26703c2948..00b364132cd 100644 --- a/src/common/snippets/src/lowered/loop_info.cpp +++ b/src/common/snippets/src/lowered/loop_info.cpp @@ -24,6 +24,10 @@ LoopInfo::LoopInfo(size_t work_amount, size_t increment, const std::vector LoopInfo::clone_loop_ports(const ExpressionMap& expr_map, return cloned_port_points; } +bool UnifiedLoopInfo::LoopPortDesc::is_dynamic() const { + return utils::is_dynamic_value(ptr_increment) || utils::is_dynamic_value(finalization_offset); +} + UnifiedLoopInfo::UnifiedLoopInfo(size_t work_amount, size_t increment, const std::vector& entries, const std::vector& exits, const SpecificIterationHandlers& handlers) @@ -173,6 +181,12 @@ std::shared_ptr UnifiedLoopInfo::clone_with_new_expr(const ExpressionM m_input_port_descs, m_output_port_descs, m_handlers); } +bool UnifiedLoopInfo::is_dynamic() const { + return LoopInfo::is_dynamic() || + std::any_of(m_input_port_descs.cbegin(), m_input_port_descs.cend(), [](const LoopPortDesc& shift) { return shift.is_dynamic(); }) || + std::any_of(m_output_port_descs.cbegin(), m_output_port_descs.cend(), [](const LoopPortDesc& shift) { return shift.is_dynamic(); }); +} + const SpecificIterationHandlers& UnifiedLoopInfo::get_handlers() const { return m_handlers; } @@ -307,7 +321,6 @@ void ExpandedLoopInfo::validate() const { "Incompatible data ptr shifts!"); } - std::shared_ptr ExpandedLoopInfo::clone_with_new_expr(const ExpressionMap& expr_map) const { const auto& new_input_ports = clone_loop_ports(expr_map, m_input_ports); const auto& new_output_ports = clone_loop_ports(expr_map, m_output_ports); @@ -316,6 +329,12 @@ std::shared_ptr ExpandedLoopInfo::clone_with_new_expr(const Expression m_ptr_increments, m_finalization_offsets, m_data_sizes, m_type, m_unified_loop_info); } +bool ExpandedLoopInfo::is_dynamic() const { + return LoopInfo::is_dynamic() || + std::any_of(m_ptr_increments.cbegin(), m_ptr_increments.cend(), [](size_t v) { return utils::is_dynamic_value(v); }) || + std::any_of(m_finalization_offsets.cbegin(), m_finalization_offsets.cend(), [](size_t v) { return utils::is_dynamic_value(v); }); +} + const std::shared_ptr& ExpandedLoopInfo::get_unified_loop_info() const { OPENVINO_ASSERT(m_unified_loop_info, "Failed to get unified loop info: it's nullptr"); return m_unified_loop_info; diff --git a/src/common/snippets/src/lowered/pass/clean_repeated_ptr_shifts.cpp b/src/common/snippets/src/lowered/pass/clean_repeated_ptr_shifts.cpp index e8aa00c426e..9552cbfdfbe 100644 --- a/src/common/snippets/src/lowered/pass/clean_repeated_ptr_shifts.cpp +++ b/src/common/snippets/src/lowered/pass/clean_repeated_ptr_shifts.cpp @@ -82,22 +82,16 @@ bool CleanRepeatedDataPointerShifts::reuse_increments(const LoopManagerPtr& loop // TODO [133463]: We have to update LoopEnd and LoopInfo since the both entities must be valid. // To avoid the both changes, we have to insert Loop ops to LinearIR in the end of pipeline. auto new_is_incremented = loop_end->get_is_incremented(); - if (const auto loop_end_dynamic = ov::as_type_ptr(loop_end_expr->get_node())) { - for (auto idx_to_drop : resetting_data_indexes) { - new_is_incremented[idx_to_drop] = false; - } - } else if (const auto loop_end_static = ov::as_type_ptr(loop_end_expr->get_node())) { - auto new_ptr_increments = loop_end_static->get_ptr_increments(); - auto new_finalization_offsets = loop_end_static->get_finalization_offsets(); - for (auto idx_to_drop : resetting_data_indexes) { - new_ptr_increments[idx_to_drop] = 0; - new_finalization_offsets[idx_to_drop] = 0; - new_is_incremented[idx_to_drop] = false; - } - loop_end_static->set_ptr_increments(new_ptr_increments); - loop_end_static->set_finalization_offsets(new_finalization_offsets); + auto new_ptr_increments = loop_end->get_ptr_increments(); + auto new_finalization_offsets = loop_end->get_finalization_offsets(); + for (auto idx_to_drop : resetting_data_indexes) { + new_is_incremented[idx_to_drop] = false; + new_ptr_increments[idx_to_drop] = 0; + new_finalization_offsets[idx_to_drop] = 0; } loop_end->set_is_incremented(new_is_incremented); + loop_end->set_ptr_increments(new_ptr_increments); + loop_end->set_finalization_offsets(new_finalization_offsets); const auto loop_info = loop_manager->get_loop_info(loop_end->get_id()); size_t loop_port_idx = 0; diff --git a/src/common/snippets/src/lowered/pass/cleanup_loop_offsets.cpp b/src/common/snippets/src/lowered/pass/cleanup_loop_offsets.cpp index 2b5162d8972..f3c34212072 100644 --- a/src/common/snippets/src/lowered/pass/cleanup_loop_offsets.cpp +++ b/src/common/snippets/src/lowered/pass/cleanup_loop_offsets.cpp @@ -5,7 +5,8 @@ #include "snippets/lowered/pass/cleanup_loop_offsets.hpp" #include "snippets/lowered/linear_ir.hpp" -#include "snippets/snippets_isa.hpp" +#include "snippets/op/loop.hpp" +#include "snippets/utils.hpp" #include "snippets/itt.hpp" namespace ov { @@ -18,7 +19,7 @@ bool CleanupLoopOffsets::run(lowered::LinearIR& linear_ir, lowered::LinearIR::co bool is_modified = false; for (auto expr_it = begin; expr_it != end; expr_it++) { const auto& node = expr_it->get()->get_node(); - if (auto loop_end = as_type_ptr(node)) { + if (auto loop_end = as_type_ptr(node)) { auto next_expr_it = std::next(expr_it); const auto& next_node = next_expr_it->get()->get_node(); // Note: Finalization offsets before the Result can be safely disregarded @@ -29,7 +30,7 @@ bool CleanupLoopOffsets::run(lowered::LinearIR& linear_ir, lowered::LinearIR::co loop_end->set_finalization_offsets(std::vector(fin_offsets.size(), 0)); is_modified = true; } - if (auto outer_loop_end = as_type_ptr(next_node)) { + if (auto outer_loop_end = as_type_ptr(next_node)) { const auto& is_incremented = loop_end->get_is_incremented(); const auto& data_sizes = loop_end->get_element_type_sizes(); auto fin_offsets = loop_end->get_finalization_offsets(); @@ -51,6 +52,8 @@ bool CleanupLoopOffsets::run(lowered::LinearIR& linear_ir, lowered::LinearIR::co if (found != per_port_connector_offset.end()) { if (!is_incremented[found->second] || outer_data_sizes[i] != data_sizes[found->second]) continue; + if (utils::is_dynamic_value(outer_ptr_increments[i]) || utils::is_dynamic_value(fin_offsets[found->second])) + continue; // Since data ptr is incremented on [ptr_increment x increment], // we should guarantee proportionality of ptr shifts. // If the data ptr can't be proportionally shifted, the optimization is not applied diff --git a/src/common/snippets/src/lowered/pass/define_buffer_clusters.cpp b/src/common/snippets/src/lowered/pass/define_buffer_clusters.cpp index 22bfe21c338..d093085dcc8 100644 --- a/src/common/snippets/src/lowered/pass/define_buffer_clusters.cpp +++ b/src/common/snippets/src/lowered/pass/define_buffer_clusters.cpp @@ -6,6 +6,7 @@ #include "snippets/lowered/pass/identify_buffers.hpp" #include "snippets/pass/tokenization.hpp" +#include "snippets/utils.hpp" #include "snippets/itt.hpp" namespace ov { @@ -46,7 +47,7 @@ size_t DefineBufferClusters::get_cluster_buffer_id(const AllocateBuffers::Buffer DefineBufferClusters::BufferPorts DefineBufferClusters::get_input_buffers(const ExpressionPtr& loop_expr) const { BufferPorts input_buffers; - const auto loop_end = ov::as_type_ptr(loop_expr->get_node()); + const auto loop_end = ov::as_type_ptr(loop_expr->get_node()); const auto in_count = loop_end->get_input_num(); const auto& connectors = loop_expr->get_input_port_connectors(); @@ -66,7 +67,7 @@ DefineBufferClusters::BufferPorts DefineBufferClusters::get_input_buffers(const DefineBufferClusters::BufferPorts DefineBufferClusters::get_output_buffers(const ExpressionPtr& loop_expr) const { BufferPorts output_buffers; - const auto loop_end = ov::as_type_ptr(loop_expr->get_node()); + const auto loop_end = ov::as_type_ptr(loop_expr->get_node()); const auto in_count = loop_end->get_input_num(); const auto out_count = loop_end->get_output_num(); const auto& connectors = loop_expr->get_input_port_connectors(); @@ -85,7 +86,7 @@ DefineBufferClusters::BufferPorts DefineBufferClusters::get_output_buffers(const void DefineBufferClusters::parse_loop(const LinearIR::constExprIt& expr_it) { const auto& expr = *expr_it; - const auto loop_end = ov::as_type_ptr(expr->get_node()); + const auto loop_end = ov::as_type_ptr(expr->get_node()); const auto& ptr_increments = loop_end->get_ptr_increments(); const auto& final_offsets = loop_end->get_finalization_offsets(); const auto& data_sizes = loop_end->get_element_type_sizes(); @@ -110,19 +111,30 @@ void DefineBufferClusters::parse_loop(const LinearIR::constExprIt& expr_it) { continue; const auto input_buffer = ov::as_type_ptr(input_buffer_expr->get_node()); + + // If allocated sizes of buffers are unkown on compilation stage (dynamic), + // we cannot be sure that they're will be the same in runtime. + if ((utils::is_dynamic_value(input_buffer->get_byte_size()) || utils::is_dynamic_value(output_buffer->get_byte_size()))) + continue; + + // Memory can be reused if reading and writing are executed proportionally: + // - the same reading/writing order + // - the same buffer memory sizes + if ((input_buffer->get_byte_size() != output_buffer->get_byte_size()) || + (input_buffer_expr->get_output_port_descriptor(0)->get_layout() != output_buffer_expr->get_input_port_descriptor(0)->get_layout())) + continue; + + // Also memory can be reused if there are the same ShiftPtrParams (data size, final offsets, ptr increments) const auto& input_buffer_ports = in.second; for (const auto& input_buffer_port_idx : input_buffer_ports) { - // Memory can be reused if reading and writing are executed proportionally: - // - the same ShiftPtrParams (data size, final offsets, ptr increments) - // - the same reading/writing order - // - the same buffer memory sizes const auto input_params = ShiftPtrParams(data_sizes[input_buffer_port_idx], ptr_increments[input_buffer_port_idx], final_offsets[input_buffer_port_idx]); const auto output_params = ShiftPtrParams(data_sizes[output_buffer_port_idx], ptr_increments[output_buffer_port_idx], final_offsets[output_buffer_port_idx]); - if (input_buffer->get_byte_size() == output_buffer->get_byte_size() && - input_buffer_expr->get_output_port_descriptor(0)->get_layout() == output_buffer_expr->get_input_port_descriptor(0)->get_layout() && - input_params == output_params) { + + // If data pointer shift parameters are unknown on model compilation stage (dynamic), + // we cannot be sure that these data pointers will be proportionally shifted in runtime. + if (input_params.is_static() && output_params.is_static() && input_params == output_params) { const auto cluster_it = find_cluster_by_expr(input_buffer_expr); OPENVINO_ASSERT(cluster_it != m_clusters.end(), "Buffer on inputs of Loop must be already saved in clusters"); // Add to the existing cluster @@ -157,11 +169,15 @@ void DefineBufferClusters::parse_nested_loops(const BufferPorts& input_buffers, auto can_be_data_ptr_proportionally_shifted = [](int64_t outer_buffer_ptr_increment, int64_t outer_buffer_data_size, int64_t inner_buffer_final_offsets, int64_t inner_buffer_data_size) { + // If data pointer shift parameters are unknown on model compilation stage (dynamic), + // we cannot be sure that these data pointers will be proportionally shifted in runtime. + if (utils::is_dynamic_value(outer_buffer_ptr_increment) || utils::is_dynamic_value(inner_buffer_final_offsets)) + return false; return (outer_buffer_ptr_increment != 0) && ((inner_buffer_data_size * inner_buffer_final_offsets * -1) == outer_buffer_ptr_increment * outer_buffer_data_size); }; - const auto outer_loop_end = ov::as_type_ptr(outer_loop_end_expr_it->get()->get_node()); + const auto outer_loop_end = ov::as_type_ptr(outer_loop_end_expr_it->get()->get_node()); const auto outer_loop_begin = outer_loop_end->get_loop_begin(); const auto& outer_ptr_increments = outer_loop_end->get_ptr_increments(); const auto& outer_data_sizes = outer_loop_end->get_element_type_sizes(); @@ -218,7 +234,7 @@ int64_t DefineBufferClusters::get_buffer_finalization_offset(const ExpressionPtr const auto consumers = buffer_out->get_consumers(); for (const auto& consumer : consumers) { const auto consumer_expr = consumer.get_expr(); - const auto loop_end = ov::as_type_ptr(consumer_expr->get_node()); + const auto loop_end = ov::as_type_ptr(consumer_expr->get_node()); if (loop_end && consumer_expr->get_loop_ids() == buffer_expr->get_loop_ids()) { const auto loop_order = ov::snippets::pass::GetTopologicalOrder(loop_end); if (loop_order > last_loop_exec_order) { @@ -243,7 +259,7 @@ bool DefineBufferClusters::unite_nested_clusters(const AllocateBuffers::BufferCl auto& up_idx = is_outer_up ? outer_idx : inner_idx; auto& down_idx = is_outer_up ? inner_idx : outer_idx; if (are_buffer_neighbours(up_buffer, down_buffer, common_loop_end_expr, up_idx, down_idx)) { - const auto common_loop_end = ov::as_type_ptr(common_loop_end_expr->get_node()); + const auto common_loop_end = ov::as_type_ptr(common_loop_end_expr->get_node()); const auto& inner_ptr_increments = common_loop_end->get_ptr_increments(); const auto& inner_final_offsets = common_loop_end->get_finalization_offsets(); const auto& inner_data_sizes = common_loop_end->get_element_type_sizes(); @@ -289,7 +305,7 @@ bool DefineBufferClusters::are_buffer_neighbours(const ExpressionPtr& up, const for (const auto& out : up->get_output_port_connectors()) { for (const auto& buffer_consumer : out->get_consumers()) { const auto buffer_consumer_expr = buffer_consumer.get_expr(); - const auto loop_end = ov::as_type_ptr(buffer_consumer_expr->get_node()); + const auto loop_end = ov::as_type_ptr(buffer_consumer_expr->get_node()); if (!loop_end) continue; const auto& loop_inputs = buffer_consumer_expr->get_input_port_connectors(); @@ -326,7 +342,7 @@ bool DefineBufferClusters::run(lowered::LinearIR& linear_ir, lowered::LinearIR:: for (auto expr_it = begin; expr_it != end; ++expr_it) { const auto& expr = *expr_it; const auto op = expr->get_node(); - if (ov::is_type(op)) { + if (ov::is_type(op)) { parse_loop(expr_it); continue; } diff --git a/src/common/snippets/src/lowered/pass/identify_buffers.cpp b/src/common/snippets/src/lowered/pass/identify_buffers.cpp index d01c0c1d2e3..7e859ce8b1b 100644 --- a/src/common/snippets/src/lowered/pass/identify_buffers.cpp +++ b/src/common/snippets/src/lowered/pass/identify_buffers.cpp @@ -4,10 +4,9 @@ #include "snippets/lowered/pass/identify_buffers.hpp" -#include "snippets/itt.hpp" #include "snippets/lowered/linear_ir.hpp" -#include "snippets/op/brgemm.hpp" #include "snippets/snippets_isa.hpp" +#include "snippets/itt.hpp" namespace ov { namespace snippets { @@ -36,9 +35,13 @@ size_t IdentifyBuffers::get_buffer_idx(const ExpressionPtr& target, const Buffer } bool IdentifyBuffers::can_reuse_id(const ShiftPtrParams& lhs, const ShiftPtrParams& rhs) { + // If data pointer shift parameters are unknown on model compilation stage (dynamic), + // we cannot be sure that these data pointers will be proportionally shifted. + // Then we force `false` value here to set unique registers for these buffers + const auto are_static = lhs.is_static() && rhs.is_static(); const auto equal_ptr_params_shifting = lhs.ptr_increment == rhs.ptr_increment && lhs.finalization_offset == rhs.finalization_offset; const auto equal_element_type_sizes = lhs.data_size == rhs.data_size; - return equal_ptr_params_shifting && (equal_element_type_sizes || (lhs.ptr_increment == 0 && lhs.finalization_offset == 0)); + return are_static && equal_ptr_params_shifting && (equal_element_type_sizes || (lhs.ptr_increment == 0 && lhs.finalization_offset == 0)); } bool IdentifyBuffers::are_adjacent(const std::pair& lhs, @@ -57,7 +60,7 @@ bool IdentifyBuffers::are_adjacent(const std::pair IdentifyBuffers::create_adjacency_matrix(LinearIR::constExprIt for (auto expr_it = begin; expr_it != end; expr_it++) { const auto &expr = *expr_it; - if (!ov::is_type(expr->get_node())) + if (!ov::is_type(expr->get_node())) continue; const auto buffer_loop_neighbours = get_buffer_loop_neighbours(expr); @@ -111,7 +114,7 @@ std::vector IdentifyBuffers::create_adjacency_matrix(LinearIR::constExprIt } IdentifyBuffers::BufferMap IdentifyBuffers::get_buffer_loop_neighbours(const ExpressionPtr& loop_end_expr) { - const auto& loop_end = ov::as_type_ptr(loop_end_expr->get_node()); + const auto& loop_end = ov::as_type_ptr(loop_end_expr->get_node()); const auto input_count = loop_end->get_input_num(); const auto output_count = loop_end->get_output_num(); @@ -142,7 +145,7 @@ IdentifyBuffers::BufferMap IdentifyBuffers::get_buffer_loop_neighbours(const Exp if (ov::is_type(child_expr->get_node())) { buffer_neighbours[child_expr] = { data_sizes[i], ptr_increments[i], finalization_offsets[i] }; buffer_count++; - } else if (ov::is_type(child_expr->get_node())) { + } else if (ov::is_type(child_expr->get_node())) { loop_count++; } } @@ -155,7 +158,7 @@ IdentifyBuffers::BufferMap IdentifyBuffers::get_buffer_loop_neighbours(const Exp } IdentifyBuffers::BufferMap IdentifyBuffers::get_buffer_loop_inside(const LinearIR::constExprIt& loop_end_it) { - const auto& loop_end = ov::as_type_ptr((*loop_end_it)->get_node()); + const auto& loop_end = ov::as_type_ptr((*loop_end_it)->get_node()); const auto loop_begin = loop_end->get_loop_begin(); BufferMap inner_buffers; for (auto it = std::reverse_iterator(loop_end_it); (*it)->get_node() != loop_begin; ++it) { diff --git a/src/common/snippets/src/lowered/pass/init_loops.cpp b/src/common/snippets/src/lowered/pass/init_loops.cpp index e9253901cfc..ec8f5a228a2 100644 --- a/src/common/snippets/src/lowered/pass/init_loops.cpp +++ b/src/common/snippets/src/lowered/pass/init_loops.cpp @@ -72,7 +72,7 @@ inline void init_is_incremented(LoopPort& port, size_t loop_id) { } } -inline int64_t get_ptr_increment(const LoopPort& loop_port, size_t work_amount) { +inline int64_t get_ptr_increment(const LoopPort& loop_port, size_t work_amount, size_t port_count) { if (!loop_port.is_incremented) return 0; @@ -87,8 +87,8 @@ inline int64_t get_ptr_increment(const LoopPort& loop_port, size_t work_amount) } else { OPENVINO_THROW("Unsupported expression port type!"); } - // When we cannot say about broadcasting by last dim - if (dim == shape.size() - 1 && utils::is_dynamic_value(shape.back())) { + // When we cannot say about broadcasting + if (utils::is_dynamic_value(shape[dim]) && port_count > 1) { return utils::get_dynamic_value(); } else if (!(shape[dim] == 1 && work_amount != 1)) { return get_stride(dim, shape); @@ -134,9 +134,12 @@ void InitLoops::init_loop_info(const UnifiedLoopInfoPtr& loop_info, const size_t init_work_amount(loop_info); const auto work_amount = loop_info->get_work_amount(); + const auto input_count = loop_info->get_input_count(); + const auto output_count = loop_info->get_output_count(); - auto init_runtime_parameters = [&work_amount](LoopPort& loop_port, UnifiedLoopInfo::LoopPortDesc& ptr_shifts_params) { - ptr_shifts_params.ptr_increment = get_ptr_increment(loop_port, work_amount); + auto init_runtime_parameters = [&work_amount, &input_count, &output_count](LoopPort& loop_port, UnifiedLoopInfo::LoopPortDesc& ptr_shifts_params) { + ptr_shifts_params.ptr_increment = get_ptr_increment(loop_port, work_amount, + loop_port.expr_port->get_type() == ExpressionPort::Input ? input_count : output_count); ptr_shifts_params.finalization_offset = get_finalization_offset(work_amount, ptr_shifts_params.ptr_increment); }; diff --git a/src/common/snippets/src/lowered/pass/insert_loops.cpp b/src/common/snippets/src/lowered/pass/insert_loops.cpp index a1079962625..07574d214de 100644 --- a/src/common/snippets/src/lowered/pass/insert_loops.cpp +++ b/src/common/snippets/src/lowered/pass/insert_loops.cpp @@ -17,41 +17,27 @@ namespace pass { void InsertLoops::insertion(LinearIR& linear_ir, const LoopManagerPtr& loop_manager, size_t loop_id) { const auto loop_info = loop_manager->get_loop_info(loop_id); - auto loop_entries = loop_info->get_input_ports(); - auto loop_exits = loop_info->get_output_ports(); const auto work_amount = loop_info->get_work_amount(); const auto work_amount_increment = loop_info->get_increment(); - - const auto loop_bounds = loop_manager->get_loop_bounds(linear_ir, loop_id); + const auto in_num = loop_info->get_input_count(); + const auto out_num = loop_info->get_output_count(); std::vector loop_end_inputs; - loop_end_inputs.reserve(loop_entries.size() + loop_exits.size()); + loop_end_inputs.reserve(in_num + out_num); loop_info->iterate_through_ports([&loop_end_inputs](const LoopPort& port) { loop_end_inputs.push_back(port.expr_port->get_port_connector_ptr()); }); const auto is_incremented = loop_info->get_is_incremented(); + const auto ptr_increments = loop_info->get_ptr_increments(); + const auto finalization_offsets = loop_info->get_finalization_offsets(); const auto io_data_sizes = loop_info->get_data_sizes(); - // Should be inited by LoopInfo - const auto is_dynamic_loop = is_loop_dynamic(loop_info); - - std::shared_ptr loop_begin = nullptr; - std::shared_ptr loop_end = nullptr; - if (is_dynamic_loop) { - loop_begin = std::make_shared(); - loop_end = std::make_shared(loop_begin, work_amount_increment, is_incremented, io_data_sizes, - loop_entries.size(), loop_exits.size(), loop_id); - - } else { - const auto ptr_increments = loop_info->get_ptr_increments(); - const auto finalization_offsets = loop_info->get_finalization_offsets(); - - loop_begin = std::make_shared(); - loop_end = std::make_shared(loop_begin, work_amount, work_amount_increment, is_incremented, ptr_increments, - finalization_offsets, io_data_sizes, loop_entries.size(), loop_exits.size(), loop_id); - } + const auto loop_begin = std::make_shared(); + const auto loop_end = std::make_shared(loop_begin, work_amount, work_amount_increment, is_incremented, ptr_increments, + finalization_offsets, io_data_sizes, in_num, out_num, loop_id); + const auto loop_bounds = loop_manager->get_loop_bounds(linear_ir, loop_id); const auto outer_loop_ids = loop_manager->get_outer_expr_loops(*loop_bounds.first, loop_id); const auto loop_begin_expr = *linear_ir.insert_node(loop_begin, std::vector{}, outer_loop_ids, false, loop_bounds.first); @@ -60,17 +46,6 @@ void InsertLoops::insertion(LinearIR& linear_ir, const LoopManagerPtr& loop_mana linear_ir.insert_node(loop_end, loop_end_inputs, outer_loop_ids, false, loop_bounds.second); } -bool InsertLoops::is_loop_dynamic(const UnifiedLoopInfoPtr& loop_info) { - auto is_loop_port_dynamic = [](const UnifiedLoopInfo::LoopPortDesc& shifts) { - return utils::is_dynamic_value(shifts.ptr_increment) || utils::is_dynamic_value(shifts.finalization_offset); - }; - const auto& entry_shifts = loop_info->get_input_port_descs(); - const auto& exit_shifts = loop_info->get_output_port_descs(); - return utils::is_dynamic_value(loop_info->get_work_amount()) || - std::any_of(entry_shifts.cbegin(), entry_shifts.cend(), is_loop_port_dynamic) || - std::any_of(exit_shifts.cbegin(), exit_shifts.cend(), is_loop_port_dynamic); -} - bool InsertLoops::run(LinearIR& linear_ir, lowered::LinearIR::constExprIt begin, lowered::LinearIR::constExprIt end) { OV_ITT_SCOPED_TASK(ov::pass::itt::domains::SnippetsTransform, "Snippets::InsertLoops") const auto& loop_manager = linear_ir.get_loop_manager(); diff --git a/src/common/snippets/src/lowered/pass/insert_specific_iterations.cpp b/src/common/snippets/src/lowered/pass/insert_specific_iterations.cpp index 3e6f6e85e0b..e89c711627a 100644 --- a/src/common/snippets/src/lowered/pass/insert_specific_iterations.cpp +++ b/src/common/snippets/src/lowered/pass/insert_specific_iterations.cpp @@ -101,12 +101,10 @@ void InsertSpecificIterations::init_decomposed_loop(LinearIR& linear_ir, LinearI const auto& loop_manager = linear_ir.get_loop_manager(); const auto new_id = loop_manager->replace_with_new_loop(linear_ir, begin, std::next(end), decomposed_loop_info, unified_loop_id); decomposed_loop_end->set_id(new_id); + decomposed_loop_end->set_work_amount(decomposed_loop_info->get_work_amount()); decomposed_loop_end->set_increment(decomposed_loop_info->get_increment()); - if (const auto static_loop_end = ov::as_type_ptr(decomposed_loop_end)) { - static_loop_end->set_work_amount(decomposed_loop_info->get_work_amount()); - static_loop_end->set_ptr_increments(decomposed_loop_info->get_ptr_increments()); - static_loop_end->set_finalization_offsets(decomposed_loop_info->get_finalization_offsets()); - } + decomposed_loop_end->set_ptr_increments(decomposed_loop_info->get_ptr_increments()); + decomposed_loop_end->set_finalization_offsets(decomposed_loop_info->get_finalization_offsets()); // Note: handlers must be run on the range started with the first operation in the loop body. const auto handlers = decomposed_loop_info->get_handler_passes(); handlers.run(linear_ir, std::next(begin), end); diff --git a/src/common/snippets/src/lowered/pass/iter_handler.cpp b/src/common/snippets/src/lowered/pass/iter_handler.cpp index 5445c229571..dd2d601366c 100644 --- a/src/common/snippets/src/lowered/pass/iter_handler.cpp +++ b/src/common/snippets/src/lowered/pass/iter_handler.cpp @@ -85,7 +85,7 @@ TransformInnerSplitLoop::TransformInnerSplitLoop(size_t tail_size) : RangedPass( bool TransformInnerSplitLoop::run(LinearIR& linear_ir, LinearIR::constExprIt begin, LinearIR::constExprIt end) { const auto& expr = *end; const auto node = expr->get_node(); - const auto loop_end = ov::as_type_ptr(node); + const auto loop_end = ov::as_type_ptr(node); OPENVINO_ASSERT(loop_end, "the last operation in range must be LoopEnd"); const auto& loop_manager = linear_ir.get_loop_manager(); @@ -97,7 +97,7 @@ bool TransformInnerSplitLoop::run(LinearIR& linear_ir, LinearIR::constExprIt beg bool modified = false; for (auto it = begin; it != end; ++it) { const auto& expr = *it; - const auto inner_loop_end = ov::as_type_ptr(expr->get_node()); + const auto inner_loop_end = ov::as_type_ptr(expr->get_node()); if (!inner_loop_end) continue; // There is already ExpandedLoopInfo @@ -105,6 +105,8 @@ bool TransformInnerSplitLoop::run(LinearIR& linear_ir, LinearIR::constExprIt beg const auto inner_dim_idx = inner_loop_info->get_dim_idx(); if (inner_dim_idx != current_dim_idx) continue; + // TODO [141735] : At the moment Splitted loops are not supported in dynamic case + OPENVINO_ASSERT(!inner_loop_end->has_dynamic_params(), "inner loop must be static in TransformInnerSplitLoop"); const auto inner_loop_begin = inner_loop_end->get_loop_begin(); const auto inner_loop_work_amount = static_cast(inner_loop_end->get_work_amount()); const auto inner_loop_increment = inner_loop_end->get_increment(); diff --git a/src/common/snippets/src/lowered/pass/optimize_loop_single_evaluation.cpp b/src/common/snippets/src/lowered/pass/optimize_loop_single_evaluation.cpp index 45c01c6644c..76921788bfd 100644 --- a/src/common/snippets/src/lowered/pass/optimize_loop_single_evaluation.cpp +++ b/src/common/snippets/src/lowered/pass/optimize_loop_single_evaluation.cpp @@ -5,7 +5,8 @@ #include "snippets/lowered/pass/optimize_loop_single_evaluation.hpp" #include "snippets/lowered/linear_ir.hpp" -#include "snippets/snippets_isa.hpp" +#include "snippets/op/loop.hpp" +#include "snippets/utils.hpp" #include "snippets/itt.hpp" namespace ov { @@ -18,7 +19,7 @@ bool OptimizeLoopSingleEvaluation::run(lowered::LinearIR& linear_ir, lowered::Li bool is_modified = false; for (auto expr_it = begin; expr_it != end; ++expr_it) { const auto& expr = *expr_it; - if (auto loop_end = ov::as_type_ptr(expr->get_node())) { + if (auto loop_end = ov::as_type_ptr(expr->get_node())) { // *1* solo vector/tail loop + empty outer loop // => skip increments (both counter & ptr) : set evaluate_once flag // *2* solo vector/tail loop + non-empty outer loop @@ -26,7 +27,7 @@ bool OptimizeLoopSingleEvaluation::run(lowered::LinearIR& linear_ir, lowered::Li // and perform pointer increments through finalization offsets // *3* vector loop(s) + one tail loop // => vector as usual, tail depends on outer loop, see *1* and *2* - if (loop_end->get_work_amount() >= 2 * loop_end->get_increment()) + if (loop_end->has_dynamic_params() || loop_end->get_work_amount() >= 2 * loop_end->get_increment()) continue; auto new_finalization_offsets = loop_end->get_finalization_offsets(); diff --git a/src/common/snippets/src/lowered/pass/split_loops.cpp b/src/common/snippets/src/lowered/pass/split_loops.cpp index 163980a21e5..dfb8edcb378 100644 --- a/src/common/snippets/src/lowered/pass/split_loops.cpp +++ b/src/common/snippets/src/lowered/pass/split_loops.cpp @@ -26,7 +26,9 @@ bool SplitLoops::can_be_split(const UnifiedLoopInfoPtr& loop_to_split, const Uni const bool equal_dim_idxes = current_dim_idx != LoopInfo::UNDEFINED_DIM_IDX && current_dim_idx == parent_dim_idx; const bool only_main_body = handlers.get_passes().empty() && handlers.get_passes().empty(); - return loop_to_split->get_work_amount() == loop_to_fuse->get_work_amount() && + // TODO [141735] : At the moment Splitted loops are not supported in dynamic case + const auto are_static = !loop_to_split->is_dynamic() && !loop_to_fuse->is_dynamic(); + return are_static && loop_to_split->get_work_amount() == loop_to_fuse->get_work_amount() && loop_to_split->get_increment() != loop_to_fuse->get_increment() && equal_dim_idxes && only_main_body; } diff --git a/src/common/snippets/src/lowered/pass/validate.cpp b/src/common/snippets/src/lowered/pass/validate.cpp index 7642e23acf8..dc9dbdea76b 100644 --- a/src/common/snippets/src/lowered/pass/validate.cpp +++ b/src/common/snippets/src/lowered/pass/validate.cpp @@ -83,23 +83,23 @@ void validate_buffer(const ExpressionPtr& expr, const LinearIR& linear_ir) { } } -void validate_loop_end_static(const ExpressionPtr& expr, const LinearIR& linear_ir) { - const auto loop_end = ov::as_type_ptr(expr->get_node()); - OPENVINO_ASSERT(loop_end, "LoopEndStatic validation expects LoopEndStatic op"); - OPENVINO_ASSERT(ov::is_type(loop_end->get_loop_begin()), - "LoopEndStatic must be connected to the LoopBeginStatic"); +void validate_loop_end(const ExpressionPtr& expr, const LinearIR& linear_ir) { + const auto loop_end = ov::as_type_ptr(expr->get_node()); + OPENVINO_ASSERT(loop_end, "LoopEnd validation expects LoopEnd op"); + OPENVINO_ASSERT(loop_end->get_loop_begin() != nullptr, + "LoopEnd must be connected to the LoopBegin"); const auto& loop_manager = linear_ir.get_loop_manager(); const auto& loop_info = loop_manager->get_loop_info(loop_end->get_id()); OPENVINO_ASSERT(loop_info->get_work_amount() == loop_end->get_work_amount() && loop_info->get_increment() == loop_end->get_increment(), - "Incompatible LoopEndStatic and the corresponding LoopInfo"); + "Incompatible LoopEnd and the corresponding LoopInfo"); const auto input_port_infos = loop_info->get_input_ports_info(); const auto output_port_infos = loop_info->get_output_ports_info(); OPENVINO_ASSERT(input_port_infos.size() == loop_end->get_input_num() && output_port_infos.size() == loop_end->get_output_num(), - "Incompatible LoopEndStatic and the corresponding LoopInfo"); + "Incompatible LoopEnd and the corresponding LoopInfo"); const auto& is_incremented = loop_end->get_is_incremented(); const auto& ptr_increments = loop_end->get_ptr_increments(); @@ -109,39 +109,12 @@ void validate_loop_end_static(const ExpressionPtr& expr, const LinearIR& linear_ OPENVINO_ASSERT(is_incremented[i + shift] == loop_port_infos[i].port.is_incremented && ptr_increments[i + shift] == loop_port_infos[i].desc.ptr_increment && final_offsets[i + shift] == loop_port_infos[i].desc.finalization_offset, - "Incompatible data ptr shifts in LoopEndStatic and the corresponding LoopInfo"); + "Incompatible data ptr shifts in LoopEnd and the corresponding LoopInfo"); } }; validate_loop_ports(input_port_infos); validate_loop_ports(output_port_infos, loop_end->get_input_num()); } - -void validate_loop_end_dynamic(const ExpressionPtr& expr, const LinearIR& linear_ir) { - const auto loop_end = ov::as_type_ptr(expr->get_node()); - OPENVINO_ASSERT(loop_end, "LoopEndDynamic validation expects LoopEndStatic op"); - OPENVINO_ASSERT(ov::is_type(loop_end->get_loop_begin()), - "LoopEndDynamic must be connected to the LoopBeginDynamic"); - - const auto& loop_manager = linear_ir.get_loop_manager(); - const auto& loop_info = loop_manager->get_loop_info(loop_end->get_id()); - OPENVINO_ASSERT(loop_info->get_increment() == loop_end->get_increment(), - "Incompatible LoopEndDynamic and the corresponding LoopInfo"); - - OPENVINO_ASSERT(loop_info->get_input_count() == loop_end->get_input_num() && - loop_info->get_output_count() == loop_end->get_output_num(), - "Incompatible LoopEndStatic and the corresponding LoopInfo"); - - const auto& is_incremented = loop_end->get_is_incremented(); - - auto validate_loop_ports = [&](const std::vector& loop_ports, size_t shift = 0) { - for (size_t i = 0; i < loop_ports.size(); ++i) { - OPENVINO_ASSERT(is_incremented[i + shift] == loop_ports[i].is_incremented, - "Incompatible data ptr shifts in LoopEndStatic and the corresponding LoopInfo"); - } - }; - validate_loop_ports(loop_info->get_input_ports()); - validate_loop_ports(loop_info->get_output_ports(), loop_end->get_input_num()); -} } // namespace Validate::Validate() { @@ -149,8 +122,7 @@ Validate::Validate() { {ov::op::v0::Parameter::get_type_info_static(), validate_parameter}, {ov::op::v0::Result::get_type_info_static(), validate_result}, {ov::snippets::op::Buffer::get_type_info_static(), validate_buffer}, - {ov::snippets::op::LoopEndStatic::get_type_info_static(), validate_loop_end_static}, - {ov::snippets::op::LoopEndDynamic::get_type_info_static(), validate_loop_end_dynamic} + {ov::snippets::op::LoopEnd::get_type_info_static(), validate_loop_end}, }; } diff --git a/src/common/snippets/src/lowered/pass/validate_expanded_loops.cpp b/src/common/snippets/src/lowered/pass/validate_expanded_loops.cpp index 2786401538c..1653d9da993 100644 --- a/src/common/snippets/src/lowered/pass/validate_expanded_loops.cpp +++ b/src/common/snippets/src/lowered/pass/validate_expanded_loops.cpp @@ -110,18 +110,16 @@ void ValidateExpandedLoops::validate_loop_expressions(const LinearIR& linear_ir) const auto expanded_loop_info = ov::as_type_ptr(loop_manager->get_loop_info(loop_id)); INFORMATIVE_ASSERT(expanded_loop_info, "expects only ExpandedLoopInfo in LoopManager"); + INFORMATIVE_ASSERT(loop_end->get_work_amount() == expanded_loop_info->get_work_amount(), + "incompatible work amount of LoopEnd and ExpandedLoopInfo"); INFORMATIVE_ASSERT(loop_end->get_increment() == expanded_loop_info->get_increment(), "incompatible increment of LoopEnd and ExpandedLoopInfo"); INFORMATIVE_ASSERT(loop_end->get_element_type_sizes() == expanded_loop_info->get_data_sizes(), "incompatible element sizes of LoopEnd and ExpandedLoopInfo"); - if (const auto static_loop_end = ov::as_type_ptr(expr->get_node())) { - INFORMATIVE_ASSERT(static_loop_end->get_work_amount() == expanded_loop_info->get_work_amount(), - "incompatible work amount of LoopEnd and ExpandedLoopInfo"); - INFORMATIVE_ASSERT(static_loop_end->get_ptr_increments() == expanded_loop_info->get_ptr_increments(), - "incompatible pointer increments of LoopEnd and ExpandedLoopInfo"); - INFORMATIVE_ASSERT(static_loop_end->get_finalization_offsets() == expanded_loop_info->get_finalization_offsets(), - "incompatible finalization offsets of LoopEnd and ExpandedLoopInfo"); - } + INFORMATIVE_ASSERT(loop_end->get_ptr_increments() == expanded_loop_info->get_ptr_increments(), + "incompatible pointer increments of LoopEnd and ExpandedLoopInfo"); + INFORMATIVE_ASSERT(loop_end->get_finalization_offsets() == expanded_loop_info->get_finalization_offsets(), + "incompatible finalization offsets of LoopEnd and ExpandedLoopInfo"); } } INFORMATIVE_ASSERT(unique_loop_ids.size() == loop_manager->get_map().size(), diff --git a/src/common/snippets/src/op/loop.cpp b/src/common/snippets/src/op/loop.cpp index 73766669300..66cdd4a275d 100644 --- a/src/common/snippets/src/op/loop.cpp +++ b/src/common/snippets/src/op/loop.cpp @@ -3,7 +3,8 @@ // #include "snippets/op/loop.hpp" -#include "snippets/generator.hpp" + +#include "snippets/utils.hpp" namespace ov { namespace snippets { @@ -29,6 +30,11 @@ void LoopBegin::validate_and_infer_types() { "LoopBegin must have LoopEnd connected to its last output"); } +std::shared_ptr LoopBegin::clone_with_new_inputs(const OutputVector& inputs) const { + OPENVINO_ASSERT(inputs.empty(), "LoopBegin should not contain inputs"); + return std::make_shared(); +} + std::shared_ptr LoopBegin::get_loop_end() const { const auto& last_output_inputs = get_output_target_inputs(0); OPENVINO_ASSERT(last_output_inputs.size() == 1, "LoopBegin has more than one inputs attached to the last output"); @@ -37,30 +43,71 @@ std::shared_ptr LoopBegin::get_loop_end() const { return loop_end; } -std::shared_ptr LoopBeginStatic::clone_with_new_inputs(const OutputVector& inputs) const { - return std::make_shared(); -} - -std::shared_ptr LoopBeginDynamic::clone_with_new_inputs(const OutputVector& inputs) const { - return std::make_shared(); -} - -LoopEnd::LoopEnd(const Output& loop_begin, size_t work_amount_increment, std::vector is_incremented, +LoopEnd::LoopEnd(const Output& loop_begin, size_t work_amount, size_t work_amount_increment, + std::vector is_incremented, std::vector ptr_increments, std::vector finalization_offsets, std::vector element_type_sizes, size_t input_num, size_t output_num, size_t id) : LoopBase({loop_begin}), m_is_incremented(std::move(is_incremented)), + m_ptr_increments(std::move(ptr_increments)), + m_finalization_offsets(std::move(finalization_offsets)), m_element_type_sizes(std::move(element_type_sizes)), + m_work_amount(work_amount), m_work_amount_increment(work_amount_increment), m_input_num(input_num), m_output_num(output_num), - m_id(id) { + m_id(id), + m_evaluate_once(false) { constructor_validate_and_infer_types(); } +void LoopEnd::validate_and_infer_types() { + NODE_VALIDATION_CHECK(this, get_input_size() == 1, "LoopEnd must have one input"); + const auto loop_begin = ov::as_type_ptr(get_input_node_shared_ptr(0)); + const auto io_size = m_input_num + m_output_num; + NODE_VALIDATION_CHECK(this, loop_begin != nullptr, "LoopEnd must have LoopBegin as the last argument"); + +#define VALIDATE_VALUES(values, name, default_value) \ + NODE_VALIDATION_CHECK(this, values.empty() || values.size() == io_size, \ + name, " must be either empty or defined per every input & output of joined Loop. Expected size: ", \ + io_size, " got ", values.size()); \ + if (values.empty()) \ + values.resize(io_size, default_value); + + VALIDATE_VALUES(m_is_incremented, "is_incremented", true) + VALIDATE_VALUES(m_ptr_increments, "ptr_increments", 0) + VALIDATE_VALUES(m_finalization_offsets, "finalization_offsets", 0) + VALIDATE_VALUES(m_element_type_sizes, "element_type_sizes", 0) +#undef VALIDATE_VALUES + + set_output_type(0, element::f32, ov::PartialShape{}); +} + +bool LoopEnd::visit_attributes(AttributeVisitor &visitor) { + std::vector int_incremented(m_is_incremented.cbegin(), m_is_incremented.cend()); + visitor.on_attribute("is_incremented", int_incremented); + visitor.on_attribute("ptr_incr", m_ptr_increments); + visitor.on_attribute("fin_offset", m_finalization_offsets); + visitor.on_attribute("data_sizes", m_element_type_sizes); + visitor.on_attribute("work_amount", m_work_amount); + visitor.on_attribute("increment", m_work_amount_increment); + visitor.on_attribute("input_num", m_input_num); + visitor.on_attribute("output_num", m_output_num); + visitor.on_attribute("id", m_id); + visitor.on_attribute("evaluate_once", m_evaluate_once); + return true; +} + +std::shared_ptr LoopEnd::clone_with_new_inputs(const OutputVector& inputs) const { + check_new_args_count(this, inputs); + const auto loop_end = std::make_shared(inputs.at(0), m_work_amount, m_work_amount_increment, m_is_incremented, m_ptr_increments, + m_finalization_offsets, m_element_type_sizes, m_input_num, m_output_num, m_id); + loop_end->m_evaluate_once = m_evaluate_once; + return loop_end; +} + std::shared_ptr LoopEnd::get_loop_begin() { const auto& loop_begin = ov::as_type_ptr(get_input_source_output(get_input_size() - 1).get_node_shared_ptr()); - if (!loop_begin) - throw std::invalid_argument("LoopEnd last input is not connected to LoopBegin"); + OPENVINO_ASSERT(loop_begin != nullptr, "LoopEnd last input is not connected to LoopBegin"); return loop_begin; } @@ -68,6 +115,14 @@ const std::vector& LoopEnd::get_is_incremented() const { return m_is_incremented; } +const std::vector& LoopEnd::get_finalization_offsets() const { + return m_finalization_offsets; +} + +const std::vector& LoopEnd::get_ptr_increments() const { + return m_ptr_increments; +} + const std::vector& LoopEnd::get_element_type_sizes() const { return m_element_type_sizes; } @@ -80,6 +135,10 @@ size_t LoopEnd::get_output_num() const { return m_output_num; } +size_t LoopEnd::get_work_amount() const { + return m_work_amount; +} + size_t LoopEnd::get_increment() const { return m_work_amount_increment; } @@ -88,131 +147,51 @@ size_t LoopEnd::get_id() const { return m_id; } +bool LoopEnd::get_evaluate_once() const { + return m_evaluate_once; +} + +bool LoopEnd::has_dynamic_params() const { + auto is_vector_dynamic = [](const std::vector& values) { + return std::any_of(values.cbegin(), values.cend(), utils::is_dynamic_value); + }; + return utils::is_dynamic_value(m_work_amount) || is_vector_dynamic(m_ptr_increments) || is_vector_dynamic(m_finalization_offsets); +} + void LoopEnd::set_is_incremented(std::vector is_incremented) { OPENVINO_ASSERT(is_incremented.size() == m_input_num + m_output_num, "LoopEnd set_is_incremented is called with inconsistent is_incremented.size()"); m_is_incremented = std::move(is_incremented); } +void LoopEnd::set_finalization_offsets(std::vector offsets) { + OPENVINO_ASSERT(offsets.size() == m_input_num + m_output_num, + "LoopEnd set_finalization_offsets is called with inconsistent offsets.size()"); + m_finalization_offsets = std::move(offsets); +} + +void LoopEnd::set_ptr_increments(std::vector new_ptr_increments) { + OPENVINO_ASSERT(new_ptr_increments.size() == m_input_num + m_output_num, + "LoopEnd set_ptr_increments is called with inconsistent new_ptr_increments.size()"); + m_ptr_increments = std::move(new_ptr_increments); +} + +void LoopEnd::set_work_amount(size_t new_work_amount) { + m_work_amount = new_work_amount; +} + void LoopEnd::set_increment(size_t new_increment) { m_work_amount_increment = new_increment; } +void LoopEnd::set_evaluate_once(bool once) { + m_evaluate_once = once; +} + void LoopEnd::set_id(size_t id) { m_id = id; } -void LoopEnd::validate_and_infer_types() { - NODE_VALIDATION_CHECK(this, get_input_size() == 1, "LoopEnd must have one input"); - const auto loop_begin = ov::as_type_ptr(get_input_node_shared_ptr(0)); - const auto io_size = m_input_num + m_output_num; - NODE_VALIDATION_CHECK(this, loop_begin != nullptr, "LoopEnd must have LoopBegin as the last argument"); - NODE_VALIDATION_CHECK(this, m_is_incremented.empty() || m_is_incremented.size() == io_size, - "is_incremented must be either empty or defined per every input & output of joined Loop. Expected size: ", - io_size, " got ", m_is_incremented.size()); - set_output_type(0, element::f32, ov::PartialShape{}); -} - -bool LoopEnd::visit_attributes(AttributeVisitor &visitor) { - std::vector int_incremented(m_is_incremented.cbegin(), m_is_incremented.cend()); - visitor.on_attribute("is_incremented", int_incremented); - visitor.on_attribute("data_sizes", m_element_type_sizes); - visitor.on_attribute("increment", m_work_amount_increment); - visitor.on_attribute("input_num", m_input_num); - visitor.on_attribute("output_num", m_output_num); - visitor.on_attribute("id", m_id); - return true; -} - -LoopEndStatic::LoopEndStatic(const Output& loop_begin, size_t work_amount, size_t work_amount_increment, - std::vector is_incremented, std::vector ptr_increments, std::vector finalization_offsets, - std::vector element_type_sizes, size_t input_num, size_t output_num, size_t id) - : LoopEnd(loop_begin, work_amount_increment, std::move(is_incremented), std::move(element_type_sizes), input_num, output_num, id), - m_ptr_increments(std::move(ptr_increments)), m_finalization_offsets(std::move(finalization_offsets)), m_work_amount(work_amount), - m_evaluate_once(false) {} - -std::shared_ptr LoopEndStatic::clone_with_new_inputs(const OutputVector& inputs) const { - check_new_args_count(this, inputs); - const auto loop_end = std::make_shared(inputs.at(0), m_work_amount, m_work_amount_increment, m_is_incremented, m_ptr_increments, - m_finalization_offsets, m_element_type_sizes, m_input_num, m_output_num, m_id); - loop_end->m_evaluate_once = m_evaluate_once; - return loop_end; -} - -void LoopEndStatic::validate_and_infer_types() { - LoopEnd::validate_and_infer_types(); - const auto io_size = m_input_num + m_output_num; - NODE_VALIDATION_CHECK(this, m_ptr_increments.empty() || m_ptr_increments.size() == io_size, - "ptr_increments must be either empty or defined per every input & output of joined Loop. Expected size: ", - io_size, " got ", m_ptr_increments.size()); - NODE_VALIDATION_CHECK(this, m_finalization_offsets.empty() || m_finalization_offsets.size() == io_size, - "finalization_offsets must be either empty or defined per every input & output of joined Loop. Expected size: ", - io_size, " got ", m_finalization_offsets.size()); - if (m_ptr_increments.empty()) - m_ptr_increments.resize(io_size, 0); - if (m_finalization_offsets.empty()) - m_finalization_offsets.resize(io_size, 0); -} - -bool LoopEndStatic::visit_attributes(AttributeVisitor &visitor) { - visitor.on_attribute("work_amount", m_work_amount); - visitor.on_attribute("ptr_incr", m_ptr_increments); - visitor.on_attribute("fin_offset", m_finalization_offsets); - visitor.on_attribute("evaluate_once", m_evaluate_once); - return LoopEnd::visit_attributes(visitor); -} - -const std::vector& LoopEndStatic::get_finalization_offsets() const { - return m_finalization_offsets; -} - -const std::vector& LoopEndStatic::get_ptr_increments() const { - return m_ptr_increments; -} - -size_t LoopEndStatic::get_work_amount() const { - return m_work_amount; -} - -bool LoopEndStatic::get_evaluate_once() const { - return m_evaluate_once; -} - -void LoopEndStatic::set_finalization_offsets(std::vector offsets) { - OPENVINO_ASSERT(offsets.size() == m_input_num + m_output_num, - "LoopEnd set_finalization_offsets is called with inconsistent offsets.size()"); - m_finalization_offsets = std::move(offsets); -} - -void LoopEndStatic::set_ptr_increments(std::vector new_ptr_increments) { - OPENVINO_ASSERT(new_ptr_increments.size() == m_input_num + m_output_num, - "LoopEnd set_ptr_increments is called with inconsistent new_ptr_increments.size()"); - m_ptr_increments = std::move(new_ptr_increments); -} -void LoopEndStatic::update_ptr_increments(int64_t new_increment) { - std::transform(m_ptr_increments.begin(), m_ptr_increments.end(), m_ptr_increments.begin(), - [new_increment](int64_t old_increment){ - return old_increment != 0 ? new_increment : 0; - }); -} - -void LoopEndStatic::set_work_amount(size_t new_work_amount) { - m_work_amount = new_work_amount; -} - -void LoopEndStatic::set_evaluate_once(bool once) { - m_evaluate_once = once; -} - -LoopEndDynamic::LoopEndDynamic(const Output& loop_begin, size_t work_amount_increment, std::vector is_incremented, - std::vector element_type_sizes, size_t input_num, size_t output_num, size_t id) - : LoopEnd(loop_begin, work_amount_increment, std::move(is_incremented), std::move(element_type_sizes), input_num, output_num, id) {} - -std::shared_ptr LoopEndDynamic::clone_with_new_inputs(const OutputVector& inputs) const { - check_new_args_count(this, inputs); - return std::make_shared(inputs.at(0), m_work_amount_increment, m_is_incremented, m_element_type_sizes, m_input_num, m_output_num, m_id); -} - } // namespace op } // namespace snippets } // namespace ov diff --git a/src/common/snippets/src/shape_inference/shape_inference.cpp b/src/common/snippets/src/shape_inference/shape_inference.cpp index d6c6081113e..ff42dae602a 100644 --- a/src/common/snippets/src/shape_inference/shape_inference.cpp +++ b/src/common/snippets/src/shape_inference/shape_inference.cpp @@ -47,12 +47,10 @@ const IShapeInferSnippetsFactory::TRegistry IShapeInferSnippetsFactory::registry SHAPE_INFER_PREDEFINED(op::HorizonMax, HorizonOpShapeInfer), SHAPE_INFER_PREDEFINED(op::HorizonSum, HorizonOpShapeInfer), // - SHAPE_INFER_PREDEFINED(op::LoopBeginStatic, SingleElementShapeInfer), - SHAPE_INFER_PREDEFINED(op::LoopBeginDynamic, SingleElementShapeInfer), + SHAPE_INFER_PREDEFINED(op::LoopBegin, SingleElementShapeInfer), SHAPE_INFER_PREDEFINED(op::Scalar, SingleElementShapeInfer), SHAPE_INFER_PREDEFINED(op::VectorBuffer, SingleElementShapeInfer), - SHAPE_INFER_PREDEFINED(op::LoopEndStatic, EmptyShapeInfer), - SHAPE_INFER_PREDEFINED(op::LoopEndDynamic, EmptyShapeInfer), + SHAPE_INFER_PREDEFINED(op::LoopEnd, EmptyShapeInfer), #ifdef SNIPPETS_DEBUG_CAPS SHAPE_INFER_PREDEFINED(op::PerfCountBegin, EmptyShapeInfer), SHAPE_INFER_PREDEFINED(op::PerfCountEnd, EmptyShapeInfer), diff --git a/src/common/snippets/tests/src/lowered/pass/loop.cpp b/src/common/snippets/tests/src/lowered/pass/loop.cpp index 560f39d96f6..0169201e0ae 100644 --- a/src/common/snippets/tests/src/lowered/pass/loop.cpp +++ b/src/common/snippets/tests/src/lowered/pass/loop.cpp @@ -79,8 +79,7 @@ static void validate(const LinearIR& linear_ir, const ref_map& reference) { size_t loop_num = 0; for (const auto& expr : linear_ir) { const auto& node = expr->get_node(); - ASSERT_TRUE(!ov::is_type(node) && !ov::is_type(node)); - const auto loop_end = ov::as_type_ptr(node); + const auto loop_end = ov::as_type_ptr(node); if (!loop_end) continue; ASSERT_GT(reference.count(loop_num), 0); diff --git a/src/common/snippets/tests/src/lowering_utils.cpp b/src/common/snippets/tests/src/lowering_utils.cpp index 35d9882e8b7..796290c3215 100644 --- a/src/common/snippets/tests/src/lowering_utils.cpp +++ b/src/common/snippets/tests/src/lowering_utils.cpp @@ -44,10 +44,8 @@ DummyTargetMachine::DummyTargetMachine(const std::vector& jitters[ov::snippets::op::BroadcastMove::get_type_info_static()] = dummy_functor; jitters[ov::snippets::op::KernelDynamic::get_type_info_static()] = dummy_functor; jitters[ov::snippets::op::KernelStatic::get_type_info_static()] = dummy_functor; - jitters[ov::snippets::op::LoopBeginDynamic::get_type_info_static()] = dummy_functor; - jitters[ov::snippets::op::LoopBeginStatic::get_type_info_static()] = dummy_functor; - jitters[ov::snippets::op::LoopEndDynamic::get_type_info_static()] = dummy_functor; - jitters[ov::snippets::op::LoopEndStatic::get_type_info_static()] = dummy_functor; + jitters[ov::snippets::op::LoopBegin::get_type_info_static()] = dummy_functor; + jitters[ov::snippets::op::LoopEnd::get_type_info_static()] = dummy_functor; #ifdef SNIPPETS_DEBUG_CAPS jitters[ov::snippets::op::PerfCountBegin::get_type_info_static()] = dummy_functor; jitters[ov::snippets::op::PerfCountEnd::get_type_info_static()] = dummy_functor; diff --git a/src/plugins/intel_cpu/src/emitters/snippets/aarch64/cpu_generator.cpp b/src/plugins/intel_cpu/src/emitters/snippets/aarch64/cpu_generator.cpp index d6b0a0fe8f5..4d9b807dee7 100644 --- a/src/plugins/intel_cpu/src/emitters/snippets/aarch64/cpu_generator.cpp +++ b/src/plugins/intel_cpu/src/emitters/snippets/aarch64/cpu_generator.cpp @@ -97,10 +97,8 @@ CPUTargetMachine::CPUTargetMachine(dnnl::impl::cpu::aarch64::cpu_isa_t host_isa) // control flow jitters[snippets::op::KernelStatic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_kernel_static_emitter); jitters[snippets::op::KernelDynamic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_kernel_dynamic_emitter); - jitters[snippets::op::LoopBeginStatic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_loop_begin_static_emitter); - jitters[snippets::op::LoopBeginDynamic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_loop_begin_dynamic_emitter); - jitters[snippets::op::LoopEndStatic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_loop_end_static_emitter); - jitters[snippets::op::LoopEndDynamic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_loop_end_dynamic_emitter); + jitters[snippets::op::LoopBegin::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_loop_begin_emitter); + jitters[snippets::op::LoopEnd::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_loop_end_emitter); // others jitters[snippets::op::Scalar::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(jit_scalar_emitter); diff --git a/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.cpp b/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.cpp index 6ca9af92254..2b5b41fb912 100644 --- a/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.cpp +++ b/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.cpp @@ -16,96 +16,39 @@ using jit_generator = dnnl::impl::cpu::aarch64::jit_generator; using cpu_isa_t = dnnl::impl::cpu::aarch64::cpu_isa_t; using ExpressionPtr = ov::snippets::lowered::ExpressionPtr; -inline static std::vector transform_idxs_to_regs(const std::vector& idxs) { - std::vector regs; - regs.resize(idxs.size(), XReg(0)); - std::transform(idxs.begin(), idxs.end(), regs.begin(), [](size_t idx){return XReg(idx);}); - return regs; -} - /* ================== jit_loop_begin_emitter ====================== */ jit_loop_begin_emitter::jit_loop_begin_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, const ov::snippets::lowered::ExpressionPtr& expr) : jit_emitter(h, isa), loop_begin_label{new Xbyak_aarch64::Label()} { - in_out_type_ = emitter_in_out_map::gpr_to_gpr; -} - -std::shared_ptr jit_loop_begin_emitter::get_loop_end(const ov::snippets::lowered::ExpressionPtr& expr) { - OV_CPU_JIT_EMITTER_ASSERT(expr->get_output_port_connectors().size() == 1, "Has invalid LoopBegin expression configuration"); - const auto& consumers = expr->get_output_port_connector(0)->get_consumers(); - OV_CPU_JIT_EMITTER_ASSERT(consumers.size() == 1, "Has invalid LoopBegin expression configuration"); - const auto loop_end = ov::as_type_ptr(consumers.cbegin()->get_expr()->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "Has invalid LoopBegin expression configuration"); - return loop_end; -} - -jit_loop_begin_static_emitter::jit_loop_begin_static_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_begin_emitter(h, isa, expr) { - OV_CPU_JIT_EMITTER_ASSERT(ov::is_type(expr->get_node()), - "Expects LoopBeginStatic expression"); - const auto loop_end = ov::as_type_ptr(get_loop_end(expr)); + const auto loop_begin = ov::as_type_ptr(expr->get_node()); + OV_CPU_JIT_EMITTER_ASSERT(loop_begin, "expects LoopBegin expression"); + const auto loop_end = loop_begin->get_loop_end(); + OV_CPU_JIT_EMITTER_ASSERT(!loop_end->has_dynamic_params(), "supports only static loops!"); work_amount = loop_end->get_work_amount(); wa_increment = loop_end->get_increment(); evaluate_once = loop_end->get_evaluate_once(); + in_out_type_ = emitter_in_out_map::gpr_to_gpr; } -void jit_loop_begin_static_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { +void jit_loop_begin_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { OV_CPU_JIT_EMITTER_ASSERT(in.empty(), "Invalid inputs size: expected 0 got " + std::to_string(in.size())); // Note: the only expected output is work amount register (communicated to jit_loop_end_emitter) OV_CPU_JIT_EMITTER_ASSERT(out.size() == 1, "Invalid outputs size: expected 1 got " + std::to_string(out.size())); + OV_CPU_JIT_EMITTER_ASSERT(loop_begin_label != nullptr, "has not inited label!"); } -void jit_loop_begin_static_emitter::emit_impl(const std::vector& in, const std::vector& out) const { - XReg reg_work_amount = XReg(out[0]); - if (!evaluate_once) { - h->mov(reg_work_amount, work_amount); - } - h->L(*loop_begin_label); -} - -void jit_loop_begin_static_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { +void jit_loop_begin_emitter::emit_code(const std::vector &in, const std::vector &out, + const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { validate_arguments(in, out); emit_impl(in, out); } -jit_loop_begin_dynamic_emitter::jit_loop_begin_dynamic_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_begin_emitter(h, isa, expr), loop_end_label(nullptr) { - OV_CPU_JIT_EMITTER_ASSERT(ov::is_type(expr->get_node()), "Expects LoopBeginDynamic expression"); - const auto loop_end = get_loop_end(expr); - wa_increment = loop_end->get_increment(); - loop_id = loop_end->get_id(); -} - -void jit_loop_begin_dynamic_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { - // Note: the only expected input is the reg_runtime_params_idx - OV_CPU_JIT_EMITTER_ASSERT(in.empty(), "Invalid inputs size: expected 0 got " + std::to_string(in.size())); - // Note: the only expected output is work amount register (communicated to jit_loop_end_emitter) - OV_CPU_JIT_EMITTER_ASSERT(out.size() == 1, "Invalid outputs size: expected 1 got " + std::to_string(out.size())); - OV_CPU_JIT_EMITTER_ASSERT(loop_end_label != nullptr && loop_begin_label != nullptr, "Has not inited labels!"); -} - -void jit_loop_begin_dynamic_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { - validate_arguments(in, out); - jit_emitter::emit_code(in, out); -} - -void jit_loop_begin_dynamic_emitter::emit_impl(const std::vector& in, const std::vector& out) const { - XReg reg_runtime_params = XReg(Operand::X0); // defined by jit_kernel_emitter +void jit_loop_begin_emitter::emit_impl(const std::vector& in, const std::vector& out) const { XReg reg_work_amount = XReg(out[0]); - XReg reg_loop_args_ptr = XReg(aux_gpr_idxs[0]); - const auto id_offset = loop_id * sizeof(jit_snippets_call_args::loop_args_t); - h->ldr(reg_loop_args_ptr, ptr(reg_runtime_params, static_cast(GET_OFF(loop_args)))); - h->ldr(reg_work_amount, ptr(reg_loop_args_ptr, static_cast(id_offset + GET_OFF_LOOP_ARGS(m_work_amount)))); - - // if wa < increment, skip the loop - h->cmp(reg_work_amount, wa_increment); - h->b(LT, *loop_end_label); - + if (!evaluate_once) { + h->mov(reg_work_amount, work_amount); + } h->L(*loop_begin_label); } @@ -118,12 +61,17 @@ jit_loop_end_emitter::jit_loop_end_emitter(dnnl::impl::cpu::aarch64::jit_generat : jit_emitter(h, isa), loop_begin_label{nullptr} { in_out_type_ = emitter_in_out_map::gpr_to_gpr; const auto loop_end = ov::as_type_ptr(expr->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "Expected LoopEnd expr"); - // Note that 1 edge connects LoopBegin and LoopEnd + OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "expected LoopEnd expr"); + OV_CPU_JIT_EMITTER_ASSERT(!loop_end->has_dynamic_params(), "supports only static loops!"); num_inputs = loop_end->get_input_num(); num_outputs = loop_end->get_output_num(); + work_amount = loop_end->get_work_amount(); wa_increment = loop_end->get_increment(); is_incremented = loop_end->get_is_incremented(); + ptr_increments = loop_end->get_ptr_increments(); + finalization_offsets = loop_end->get_finalization_offsets(); + data_sizes = loop_end->get_element_type_sizes(); + evaluate_once = loop_end->get_evaluate_once(); const auto begin_expr = get_loop_begin_expr(expr); const auto& loop_begin_emitter = std::dynamic_pointer_cast(begin_expr->get_emitter()); @@ -138,36 +86,25 @@ ov::snippets::lowered::ExpressionPtr jit_loop_end_emitter::get_loop_begin_expr(c return begin_expr; } -jit_loop_end_static_emitter::jit_loop_end_static_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_end_emitter(h, isa, expr) { - const auto loop_end = ov::as_type_ptr(expr->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "Expected LoopEndStatic expr"); - work_amount = static_cast(loop_end->get_work_amount()); - is_incremented = loop_end->get_is_incremented(); - ptr_increments = loop_end->get_ptr_increments(); - finalization_offsets = loop_end->get_finalization_offsets(); - data_sizes = loop_end->get_element_type_sizes(); - evaluate_once = loop_end->get_evaluate_once(); -} - -void jit_loop_end_static_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { - const auto io_size = num_inputs + num_outputs; +void jit_loop_end_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { +const auto io_size = num_inputs + num_outputs; OV_CPU_JIT_EMITTER_ASSERT(out.size() == 0, "Invalid number of out arguments: expected ", 0, " got ", out.size()); OV_CPU_JIT_EMITTER_ASSERT(in.size() == io_size + 1, "Invalid number of in arguments: expected ", io_size + 1, " got ", in.size()); + OV_CPU_JIT_EMITTER_ASSERT(is_incremented.size() == io_size, "Invalid is_incremented size: expected ", io_size, " got ", is_incremented.size()); OV_CPU_JIT_EMITTER_ASSERT(ptr_increments.size() == io_size, "Invalid ptr_increments size: expected ", io_size, " got ", ptr_increments.size()); OV_CPU_JIT_EMITTER_ASSERT(finalization_offsets.size() == io_size, "Invalid finalization_offsets size: expected: ", io_size, " got ", finalization_offsets.size()); OV_CPU_JIT_EMITTER_ASSERT(data_sizes.size() == io_size, "Invalid data_sizes size: expected: ", io_size, " got ", data_sizes.size()); + OV_CPU_JIT_EMITTER_ASSERT(loop_begin_label != nullptr, "has not inited begin label!"); } -void jit_loop_end_static_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { +void jit_loop_end_emitter::emit_code(const std::vector &in, const std::vector &out, + const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { validate_arguments(in, out); emit_impl(in, out); } -void jit_loop_end_static_emitter::emit_impl(const std::vector& in, const std::vector& out) const { +void jit_loop_end_emitter::emit_impl(const std::vector& in, const std::vector& out) const { std::vector data_ptr_reg_idxs; data_ptr_reg_idxs.reserve(num_inputs + num_outputs); std::copy(in.begin(), in.end() - 1, std::back_inserter(data_ptr_reg_idxs)); @@ -201,68 +138,6 @@ void jit_loop_end_static_emitter::emit_impl(const std::vector& in, const } } -jit_loop_end_dynamic_emitter::jit_loop_end_dynamic_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_end_emitter(h, isa, expr), loop_end_label{new Xbyak_aarch64::Label()} { - const auto loop_end = ov::as_type_ptr(expr->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "Expected LoopEndDynamic expr"); - loop_id = loop_end->get_id(); - - const auto begin_expr = get_loop_begin_expr(expr); - const auto& loop_begin_emitter = std::dynamic_pointer_cast(begin_expr->get_emitter()); - OV_CPU_JIT_EMITTER_ASSERT(loop_begin_emitter, "LoopBeginDynamic expected jit_loop_begin_dynamic_emitter"); - loop_begin_emitter->set_loop_end_label(loop_end_label); -} - -void jit_loop_end_dynamic_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { - OV_CPU_JIT_EMITTER_ASSERT(loop_end_label != nullptr && loop_begin_label != nullptr, "Has not inited labels!"); - // Note: there must be additional input argument for runtime parameters - const auto io_size = num_inputs + num_outputs; - OV_CPU_JIT_EMITTER_ASSERT(in.size() == io_size + 1, "Invalid number of in arguments: expected ", io_size + 1, " got ", in.size()); - OV_CPU_JIT_EMITTER_ASSERT(out.size() == 0, "Invalid number of out arguments: expected ", 0, " got ", out.size()); -} - -void jit_loop_end_dynamic_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { - validate_arguments(in, out); - jit_emitter::emit_code(in, out); -} - -void jit_loop_end_dynamic_emitter::emit_impl(const std::vector& in, const std::vector& out) const { - XReg reg_runtime_params = XReg(Operand::X0); // defined by jit_kernel_emitter - XReg reg_work_amount = XReg(in.back()); - XReg reg_increments = XReg(aux_gpr_idxs[0]); - XReg reg_aux = XReg(aux_gpr_idxs[1]); - const auto id_offset = loop_id * sizeof(jit_snippets_call_args::loop_args_t); - - std::vector data_ptr_regs = transform_idxs_to_regs(std::vector(in.begin(), in.end() - 1)); - - // todo: Note that we can pre-save reg_loop_args_ptr in jit_loop_begin_dynamic_emitter and pass it here like work_amount_reg - // this would save us one dereferencing here and in finalization offsets - h->ldr(reg_increments, ptr(reg_runtime_params, static_cast(GET_OFF(loop_args)))); - h->ldr(reg_increments, ptr(reg_increments, static_cast(id_offset + GET_OFF_LOOP_ARGS(m_ptr_increments)))); - for (size_t idx = 0; idx < data_ptr_regs.size(); idx++) { - if (is_incremented[idx]) { - h->ldr(reg_aux, ptr(reg_increments, static_cast(idx * sizeof(int64_t)))); - h->add(data_ptr_regs[idx], data_ptr_regs[idx], reg_aux); - } - } - h->sub_imm(reg_work_amount, reg_work_amount, wa_increment, h->X_TMP_0); - h->cmp(reg_work_amount, wa_increment); - h->b(GE, *loop_begin_label); - - h->ldr(reg_increments, ptr(reg_runtime_params, static_cast(GET_OFF(loop_args)))); - h->ldr(reg_increments, ptr(reg_increments, static_cast(id_offset + GET_OFF_LOOP_ARGS(m_finalization_offsets)))); - for (size_t idx = 0; idx < data_ptr_regs.size(); idx++) { - if (is_incremented[idx]) { - h->ldr(reg_aux, ptr(reg_increments, static_cast(idx * sizeof(int64_t)))); - h->add(data_ptr_regs[idx], data_ptr_regs[idx], reg_aux); - } - } - - h->L(*loop_end_label); -} - /* ============================================================== */ } // namespace aarch64 diff --git a/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.hpp b/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.hpp index af75ac1eb41..6ec87835821 100644 --- a/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.hpp +++ b/src/plugins/intel_cpu/src/emitters/snippets/aarch64/jit_loop_emitters.hpp @@ -21,49 +21,19 @@ public: size_t get_inputs_count() const override { return 0; } + void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, + const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; + std::shared_ptr get_begin_label() { return loop_begin_label; } protected: - static std::shared_ptr get_loop_end(const ov::snippets::lowered::ExpressionPtr& expr); + void validate_arguments(const std::vector &in, const std::vector &out) const override; + void emit_impl(const std::vector& in, const std::vector& out) const override; std::shared_ptr loop_begin_label; - int64_t wa_increment = 0; -}; - -class jit_loop_begin_static_emitter: public jit_loop_begin_emitter { -public: - jit_loop_begin_static_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, - const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; -protected: - void validate_arguments(const std::vector &in, const std::vector &out) const override; - void emit_impl(const std::vector& in, const std::vector& out) const override; - - bool evaluate_once = false; size_t work_amount = 0; -}; - -class jit_loop_begin_dynamic_emitter: public jit_loop_begin_emitter { -public: - jit_loop_begin_dynamic_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, - const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; - - void set_loop_end_label(const std::shared_ptr& label) { loop_end_label = label; } - -protected: - void validate_arguments(const std::vector &in, const std::vector &out) const override; - void emit_impl(const std::vector& in, const std::vector& out) const override; - - // For Loop arguments - size_t get_aux_gprs_count() const override { return 1; } - - std::shared_ptr loop_end_label; - size_t loop_id; + int64_t wa_increment = 0; + bool evaluate_once = false; }; /* ============================================================== */ @@ -77,21 +47,6 @@ public: size_t get_inputs_count() const override { return 0; } -protected: - static ov::snippets::lowered::ExpressionPtr get_loop_begin_expr(const ov::snippets::lowered::ExpressionPtr& expr); - - std::shared_ptr loop_begin_label; - size_t num_inputs = 0; - size_t num_outputs = 0; - int64_t wa_increment = 0; - std::vector is_incremented = {}; -}; - -class jit_loop_end_static_emitter: public jit_loop_end_emitter { -public: - jit_loop_end_static_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; @@ -99,7 +54,13 @@ protected: void validate_arguments(const std::vector &in, const std::vector &out) const override; void emit_impl(const std::vector& in, const std::vector& out) const override; + static ov::snippets::lowered::ExpressionPtr get_loop_begin_expr(const ov::snippets::lowered::ExpressionPtr& expr); + + std::shared_ptr loop_begin_label; + size_t num_inputs = 0; + size_t num_outputs = 0; size_t work_amount = 0; + int64_t wa_increment = 0; std::vector is_incremented = {}; std::vector ptr_increments = {}; std::vector finalization_offsets = {}; @@ -107,25 +68,6 @@ protected: bool evaluate_once = false; }; -class jit_loop_end_dynamic_emitter: public jit_loop_end_emitter { -public: - jit_loop_end_dynamic_emitter(dnnl::impl::cpu::aarch64::jit_generator* h, dnnl::impl::cpu::aarch64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, - const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; - -protected: - void validate_arguments(const std::vector &in, const std::vector &out) const override; - void emit_impl(const std::vector& in, const std::vector& out) const override; - - // For Loop arguments - size_t get_aux_gprs_count() const override { return 2; } - - std::shared_ptr loop_end_label; - size_t loop_id; -}; - /* ============================================================== */ } // namespace aarch64 diff --git a/src/plugins/intel_cpu/src/emitters/snippets/x64/cpu_generator.cpp b/src/plugins/intel_cpu/src/emitters/snippets/x64/cpu_generator.cpp index 323bae34806..a1fde3bf28f 100644 --- a/src/plugins/intel_cpu/src/emitters/snippets/x64/cpu_generator.cpp +++ b/src/plugins/intel_cpu/src/emitters/snippets/x64/cpu_generator.cpp @@ -239,10 +239,8 @@ intel_cpu::CPUTargetMachine::CPUTargetMachine(dnnl::impl::cpu::x64::cpu_isa_t ho jitters[snippets::op::KernelStatic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_kernel_static_emitter); jitters[snippets::op::KernelDynamic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_kernel_dynamic_emitter); - jitters[snippets::op::LoopBeginStatic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_loop_begin_static_emitter); - jitters[snippets::op::LoopBeginDynamic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_loop_begin_dynamic_emitter); - jitters[snippets::op::LoopEndStatic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_loop_end_static_emitter); - jitters[snippets::op::LoopEndDynamic::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_loop_end_dynamic_emitter); + jitters[snippets::op::LoopBegin::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_loop_begin_emitter); + jitters[snippets::op::LoopEnd::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_loop_end_emitter); // Note: jit_brgemm_emitter supports runtime recompilation, so its constructor takes additional arguments jitters[intel_cpu::BrgemmCPU::get_type_info_static()] = CREATE_SNIPPETS_EMITTER(intel_cpu::jit_brgemm_emitter, kernel_executor_table, diff --git a/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.cpp b/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.cpp index ee6225271f6..566e495e88d 100644 --- a/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.cpp +++ b/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.cpp @@ -5,6 +5,7 @@ #include "jit_loop_emitters.hpp" #include "emitters/snippets/jit_snippets_call_args.hpp" +#include "snippets/utils.hpp" using namespace Xbyak; using namespace dnnl::impl; @@ -13,89 +14,57 @@ using namespace dnnl::impl::cpu::x64; namespace ov { namespace intel_cpu { -inline static void transform_idxs_to_regs(const std::vector& idxs, std::vector& regs) { - regs.resize(idxs.size()); - std::transform(idxs.begin(), idxs.end(), regs.begin(), [](size_t idx){ return Xbyak::Reg64(static_cast(idx)); }); -} - /* ================== jit_loop_begin_emitter ====================== */ jit_loop_begin_emitter::jit_loop_begin_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const ov::snippets::lowered::ExpressionPtr& expr) - : jit_emitter(h, isa), loop_begin_label{new Xbyak::Label()} { - in_out_type_ = emitter_in_out_map::gpr_to_gpr; -} - -std::shared_ptr jit_loop_begin_emitter::get_loop_end(const ov::snippets::lowered::ExpressionPtr& expr) { - OV_CPU_JIT_EMITTER_ASSERT(expr->get_output_port_connectors().size() == 1, "has invalid LoopBegin expression configuration"); - const auto& consumers = expr->get_output_port_connector(0)->get_consumers(); - OV_CPU_JIT_EMITTER_ASSERT(consumers.size() == 1, "has invalid LoopBegin expression configuration"); - const auto loop_end = ov::as_type_ptr(consumers.cbegin()->get_expr()->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "has invalid LoopBegin expression configuration"); - return loop_end; -} - -jit_loop_begin_static_emitter::jit_loop_begin_static_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_begin_emitter(h, isa, expr) { - OV_CPU_JIT_EMITTER_ASSERT(ov::is_type(expr->get_node()), - "expects LoopBeginStatic expression"); - const auto loop_end = ov::as_type_ptr(get_loop_end(expr)); + : jit_emitter(h, isa), loop_begin_label{new Xbyak::Label()}, loop_end_label(nullptr) { + const auto loop_begin = ov::as_type_ptr(expr->get_node()); + OV_CPU_JIT_EMITTER_ASSERT(loop_begin, "expects LoopBegin expression"); + const auto loop_end = loop_begin->get_loop_end(); work_amount = loop_end->get_work_amount(); wa_increment = loop_end->get_increment(); evaluate_once = loop_end->get_evaluate_once(); + loop_id = loop_end->get_id(); + is_work_amount_dynamic = ov::snippets::utils::is_dynamic_value(work_amount); + in_out_type_ = emitter_in_out_map::gpr_to_gpr; } -void jit_loop_begin_static_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { +size_t jit_loop_begin_emitter::aux_gprs_count() const { + // We should have aux GPR to store Loop arguments from `runtime_args` + // where we will take all needed information about the current loop: work amount + return is_work_amount_dynamic ? 1 : 0; +} + +void jit_loop_begin_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { OV_CPU_JIT_EMITTER_ASSERT(in.empty(), "Invalid inputs size: expected 0 got " + std::to_string(in.size())); // Note: the only expected output is work amount register (communicated to jit_loop_end_emitter) OV_CPU_JIT_EMITTER_ASSERT(out.size() == 1, "Invalid outputs size: expected 1 got " + std::to_string(out.size())); + OV_CPU_JIT_EMITTER_ASSERT(loop_begin_label != nullptr && loop_end_label != nullptr, "has not inited labels!"); + OV_CPU_JIT_EMITTER_ASSERT(implication(is_work_amount_dynamic, !evaluate_once), "with dynamic work_amount cannot evaluate once!"); } -void jit_loop_begin_static_emitter::emit_impl(const std::vector& in, const std::vector& out) const { - Xbyak::Reg64 reg_work_amount = Xbyak::Reg64(static_cast(out.back())); - if (!evaluate_once) { +void jit_loop_begin_emitter::emit_code(const std::vector &in, const std::vector &out, + const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { + validate_arguments(in, out); + jit_emitter::emit_code(in, out, pool_vec_idxs, pool_gpr_idxs); +} + +void jit_loop_begin_emitter::emit_impl(const std::vector& in, const std::vector& out) const { + // If the loop evaulate once, we can skip loop begin code emission + if (evaluate_once) + return; + + Reg64 reg_work_amount = Reg64(static_cast(out.back())); + if (is_work_amount_dynamic) { + Reg64 reg_runtime_params = abi_param1; // defined by jit_kernel_emitter + Reg64 reg_loop_args_ptr = Reg64(static_cast(aux_gpr_idxs[0])); + const auto id_offset = loop_id * sizeof(jit_snippets_call_args::loop_args_t); + h->mov(reg_loop_args_ptr, h->ptr[reg_runtime_params + GET_OFF(loop_args)]); + h->mov(reg_work_amount, h->ptr[reg_loop_args_ptr + id_offset + GET_OFF_LOOP_ARGS(m_work_amount)]); + } else { h->mov(reg_work_amount, work_amount); } - h->L(*loop_begin_label); -} - -void jit_loop_begin_static_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { - validate_arguments(in, out); - emit_impl(in, out); -} - -jit_loop_begin_dynamic_emitter::jit_loop_begin_dynamic_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_begin_emitter(h, isa, expr), loop_end_label(nullptr) { - OV_CPU_JIT_EMITTER_ASSERT(ov::is_type(expr->get_node()), "expects LoopBeginDynamic expression"); - const auto loop_end = get_loop_end(expr); - wa_increment = loop_end->get_increment(); - loop_id = loop_end->get_id(); -} - -void jit_loop_begin_dynamic_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { - // Note: the only expected input is the reg_runtime_params_idx - OV_CPU_JIT_EMITTER_ASSERT(in.empty(), "Invalid inputs size: expected 0 got " + std::to_string(in.size())); - // Note: the only expected output is work amount register (communicated to jit_loop_end_emitter) - OV_CPU_JIT_EMITTER_ASSERT(out.size() == 1, "Invalid outputs size: expected 1 got " + std::to_string(out.size())); - OV_CPU_JIT_EMITTER_ASSERT(loop_end_label != nullptr && loop_begin_label != nullptr, "has not inited labels!"); -} - -void jit_loop_begin_dynamic_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { - validate_arguments(in, out); - jit_emitter::emit_code(in, out); -} - -void jit_loop_begin_dynamic_emitter::emit_impl(const std::vector& in, const std::vector& out) const { - Xbyak::Reg64 reg_runtime_params = abi_param1; // defined by jit_kernel_emitter - Xbyak::Reg64 reg_work_amount = Xbyak::Reg64(static_cast(out.back())); - Xbyak::Reg64 reg_loop_args_ptr = Xbyak::Reg64(static_cast(aux_gpr_idxs[0])); - const auto id_offset = loop_id * sizeof(jit_snippets_call_args::loop_args_t); - h->mov(reg_loop_args_ptr, h->ptr[reg_runtime_params + GET_OFF(loop_args)]); - h->mov(reg_work_amount, h->ptr[reg_loop_args_ptr + id_offset + GET_OFF_LOOP_ARGS(m_work_amount)]); // if wa < increment, skip the loop h->cmp(reg_work_amount, wa_increment); @@ -110,19 +79,31 @@ void jit_loop_begin_dynamic_emitter::emit_impl(const std::vector& in, co jit_loop_end_emitter::jit_loop_end_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, const ov::snippets::lowered::ExpressionPtr& expr) - : jit_emitter(h, isa), loop_begin_label{nullptr} { + : jit_emitter(h, isa), loop_begin_label{nullptr}, loop_end_label{new Xbyak::Label()} { in_out_type_ = emitter_in_out_map::gpr_to_gpr; const auto loop_end = ov::as_type_ptr(expr->get_node()); OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "expected LoopEnd expr"); - // Note that 1 edge connects LoopBegin and LoopEnd num_inputs = loop_end->get_input_num(); num_outputs = loop_end->get_output_num(); + work_amount = loop_end->get_work_amount(); wa_increment = loop_end->get_increment(); is_incremented = loop_end->get_is_incremented(); + ptr_increments = loop_end->get_ptr_increments(); + finalization_offsets = loop_end->get_finalization_offsets(); + data_sizes = loop_end->get_element_type_sizes(); + evaluate_once = loop_end->get_evaluate_once(); + loop_id = loop_end->get_id(); + + are_ptr_increments_dynamic = + std::any_of(ptr_increments.cbegin(), ptr_increments.cend(), ov::snippets::utils::is_dynamic_value); + are_final_offsets_dynamic = + std::any_of(finalization_offsets.cbegin(), finalization_offsets.cend(), ov::snippets::utils::is_dynamic_value); + are_ptr_shifts_dynamic = are_ptr_increments_dynamic || are_final_offsets_dynamic; const auto begin_expr = get_loop_begin_expr(expr); const auto& loop_begin_emitter = std::dynamic_pointer_cast(begin_expr->get_emitter()); OV_CPU_JIT_EMITTER_ASSERT(loop_begin_emitter, "LoopBegin expected jit_loop_begin_emitter"); + loop_begin_emitter->set_loop_end_label(loop_end_label); loop_begin_label = loop_begin_emitter->get_begin_label(); } @@ -133,116 +114,69 @@ ov::snippets::lowered::ExpressionPtr jit_loop_end_emitter::get_loop_begin_expr(c return begin_expr; } -jit_loop_end_static_emitter::jit_loop_end_static_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_end_emitter(h, isa, expr) { - const auto loop_end = ov::as_type_ptr(expr->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "expected LoopEndStatic expr"); - work_amount = static_cast(loop_end->get_work_amount()); - is_incremented = loop_end->get_is_incremented(); - ptr_increments = loop_end->get_ptr_increments(); - finalization_offsets = loop_end->get_finalization_offsets(); - data_sizes = loop_end->get_element_type_sizes(); - evaluate_once = loop_end->get_evaluate_once(); -} - -void jit_loop_end_static_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { +void jit_loop_end_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { const auto io_size = num_inputs + num_outputs; OV_CPU_JIT_EMITTER_ASSERT(out.size() == 0, "Invalid number of out arguments: expected ", 0, " got ", out.size()); OV_CPU_JIT_EMITTER_ASSERT(in.size() == io_size + 1, "Invalid number of in arguments: expected ", io_size + 1, " got ", in.size()); + OV_CPU_JIT_EMITTER_ASSERT(is_incremented.size() == io_size, "Invalid is_incremented size: expected ", io_size, " got ", is_incremented.size()); OV_CPU_JIT_EMITTER_ASSERT(ptr_increments.size() == io_size, "Invalid ptr_increments size: expected ", io_size, " got ", ptr_increments.size()); OV_CPU_JIT_EMITTER_ASSERT(finalization_offsets.size() == io_size, "Invalid finalization_offsets size: expected: ", io_size, " got ", finalization_offsets.size()); OV_CPU_JIT_EMITTER_ASSERT(data_sizes.size() == io_size, "Invalid data_sizes size: expected: ", io_size, " got ", data_sizes.size()); + OV_CPU_JIT_EMITTER_ASSERT(loop_end_label != nullptr && loop_begin_label != nullptr, "has not inited labels!"); + OV_CPU_JIT_EMITTER_ASSERT(implication(are_ptr_shifts_dynamic, !evaluate_once), "with dynamic data pointer shifts cannot evaluate once!"); } -void jit_loop_end_static_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { +void jit_loop_end_emitter::emit_code(const std::vector &in, const std::vector &out, + const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { validate_arguments(in, out); - emit_impl(in, out); + jit_emitter::emit_code(in, out, pool_vec_idxs, pool_gpr_idxs); } -void jit_loop_end_static_emitter::emit_impl(const std::vector& in, const std::vector& out) const { +size_t jit_loop_end_emitter::aux_gprs_count() const { + // We should have aux GPR to store Loop arguments from `runtime_args` + // where we will take all needed information about the current loop: data pointer shifts + return are_ptr_shifts_dynamic ? 1 : 0; +} + +void jit_loop_end_emitter::emit_impl(const std::vector& in, const std::vector& out) const { std::vector data_ptr_reg_idxs; // the last input is actually a work_amount reg - data_ptr_reg_idxs.reserve(num_inputs - 1); + data_ptr_reg_idxs.reserve(num_inputs + num_outputs); std::copy(in.begin(), in.end() - 1, std::back_inserter(data_ptr_reg_idxs)); - Reg64 reg_work_amount = Reg64(in.back()); - if (!evaluate_once) { - for (size_t idx = 0; idx < data_ptr_reg_idxs.size(); idx++) { - if (!is_incremented[idx] || ptr_increments[idx] == 0) - continue; - Reg64 data_reg = Reg64(static_cast(data_ptr_reg_idxs[idx])); - h->add(data_reg, ptr_increments[idx] * wa_increment * data_sizes[idx]); + const auto id_offset = loop_id * sizeof(jit_snippets_call_args::loop_args_t); + Reg64 reg_increments = are_ptr_shifts_dynamic ? Reg64(static_cast(aux_gpr_idxs[0])) : Reg64(); + + auto apply_increments = [&](bool use_runtime_args, size_t field_offset, const std::vector& increments, size_t scale) { + if (use_runtime_args) { + Reg64 reg_runtime_params = abi_param1; /* defined by jit_kernel_emitter */ + h->mov(reg_increments, h->ptr[reg_runtime_params + GET_OFF(loop_args)]); + h->mov(reg_increments, h->ptr[reg_increments + id_offset + field_offset]); } + for (size_t idx = 0; idx < data_ptr_reg_idxs.size(); idx++) { + const auto& increment = increments[idx]; + if (is_incremented[idx] && increment != 0) { + if (ov::snippets::utils::is_dynamic_value(increment)) { + OV_CPU_JIT_EMITTER_ASSERT(use_runtime_args, "Loop argument structure cannot be pushed to aux GPR"); + h->add(Reg64(static_cast(data_ptr_reg_idxs[idx])), h->ptr[reg_increments + idx * sizeof(int64_t)]); + } else { + h->add(Reg64(static_cast(data_ptr_reg_idxs[idx])), increment * scale * data_sizes[idx]); + } + } + } + }; + + if (!evaluate_once) { + apply_increments(are_ptr_increments_dynamic, GET_OFF_LOOP_ARGS(m_ptr_increments), ptr_increments, wa_increment); + + Reg64 reg_work_amount = Reg64(in.back()); h->sub(reg_work_amount, wa_increment); h->cmp(reg_work_amount, wa_increment); - h->jge(*loop_begin_label); + h->jge(*loop_begin_label, Xbyak::CodeGenerator::T_NEAR); } - for (size_t idx = 0; idx < data_ptr_reg_idxs.size(); idx++) { - if (!is_incremented[idx] || finalization_offsets[idx] == 0) - continue; - Reg64 data_reg = Reg64(static_cast(data_ptr_reg_idxs[idx])); - h->add(data_reg, finalization_offsets[idx] * data_sizes[idx]); - } -} - -jit_loop_end_dynamic_emitter::jit_loop_end_dynamic_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr) - : jit_loop_end_emitter(h, isa, expr), loop_end_label{new Xbyak::Label()} { - const auto loop_end = ov::as_type_ptr(expr->get_node()); - OV_CPU_JIT_EMITTER_ASSERT(loop_end != nullptr, "expected LoopEndDynamic expr"); - loop_id = loop_end->get_id(); - - const auto begin_expr = get_loop_begin_expr(expr); - const auto& loop_begin_emitter = std::dynamic_pointer_cast(begin_expr->get_emitter()); - OV_CPU_JIT_EMITTER_ASSERT(loop_begin_emitter, "LoopBeginDynamic expected jit_loop_begin_dynamic_emitter"); - loop_begin_emitter->set_loop_end_label(loop_end_label); -} - -void jit_loop_end_dynamic_emitter::validate_arguments(const std::vector &in, const std::vector &out) const { - OV_CPU_JIT_EMITTER_ASSERT(loop_end_label != nullptr && loop_begin_label != nullptr, "has not inited labels!"); - // Note: there must be additional input argument for runtime parameters - const auto io_size = num_inputs + num_outputs; - OV_CPU_JIT_EMITTER_ASSERT(in.size() == io_size + 1, "Invalid number of in arguments: expected ", io_size + 1, " got ", in.size()); - OV_CPU_JIT_EMITTER_ASSERT(out.size() == 0, "Invalid number of out arguments: expected ", 0, " got ", out.size()); -} - -void jit_loop_end_dynamic_emitter::emit_code(const std::vector &in, const std::vector &out, - const std::vector &pool_vec_idxs, const std::vector &pool_gpr_idxs) const { - validate_arguments(in, out); - jit_emitter::emit_code(in, out); -} - -void jit_loop_end_dynamic_emitter::emit_impl(const std::vector& in, const std::vector& out) const { - Xbyak::Reg64 reg_runtime_params = abi_param1; // defined by jit_kernel_emitter - Xbyak::Reg64 reg_work_amount = Xbyak::Reg64(static_cast(in[in.size() - 1])); - Xbyak::Reg64 reg_increments = Xbyak::Reg64(static_cast(aux_gpr_idxs[0])); - const auto id_offset = loop_id * sizeof(jit_snippets_call_args::loop_args_t); - - std::vector data_ptr_regs; - transform_idxs_to_regs(std::vector(in.begin(), in.end() - 1), data_ptr_regs); - - // todo: Note that we can pre-save reg_loop_args_ptr in jit_loop_begin_dynamic_emitter and pass it here like work_amount_reg - // this would save us one dereferencing here and in finalization offsets - h->mov(reg_increments, h->ptr[reg_runtime_params + GET_OFF(loop_args)]); - h->mov(reg_increments, h->ptr[reg_increments + id_offset + GET_OFF_LOOP_ARGS(m_ptr_increments)]); - for (size_t idx = 0; idx < data_ptr_regs.size(); idx++) { - if (is_incremented[idx]) - h->add(data_ptr_regs[idx], h->ptr[reg_increments + idx * sizeof(int64_t)]); - } - h->sub(reg_work_amount, wa_increment); - h->cmp(reg_work_amount, wa_increment); - h->jge(*loop_begin_label, Xbyak::CodeGenerator::T_NEAR); - - h->mov(reg_increments, h->ptr[reg_runtime_params + GET_OFF(loop_args)]); - h->mov(reg_increments, h->ptr[reg_increments + id_offset + GET_OFF_LOOP_ARGS(m_finalization_offsets)]); - for (size_t idx = 0; idx < data_ptr_regs.size(); idx++) { - if (is_incremented[idx]) - h->add(data_ptr_regs[idx], h->ptr[reg_increments + idx * sizeof(int64_t)]); - } + apply_increments(are_final_offsets_dynamic, GET_OFF_LOOP_ARGS(m_finalization_offsets), finalization_offsets, 1); h->L(*loop_end_label); } diff --git a/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.hpp b/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.hpp index 1f1013dfca7..0af5ac46621 100644 --- a/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.hpp +++ b/src/plugins/intel_cpu/src/emitters/snippets/x64/jit_loop_emitters.hpp @@ -7,6 +7,7 @@ #include "emitters/plugin/x64/jit_emitter.hpp" #include "snippets/op/loop.hpp" +#include "snippets/utils.hpp" namespace ov { namespace intel_cpu { @@ -20,50 +21,27 @@ public: size_t get_inputs_num() const override { return 0; } + void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, + const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; + + void set_loop_end_label(const std::shared_ptr& label) { loop_end_label = label; } std::shared_ptr get_begin_label() { return loop_begin_label; } -protected: - static std::shared_ptr get_loop_end(const ov::snippets::lowered::ExpressionPtr& expr); - - std::shared_ptr loop_begin_label; - int64_t wa_increment = 0; -}; - -class jit_loop_begin_static_emitter: public jit_loop_begin_emitter { -public: - jit_loop_begin_static_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, - const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; protected: void validate_arguments(const std::vector &in, const std::vector &out) const override; void emit_impl(const std::vector& in, const std::vector& out) const override; - bool evaluate_once = false; + size_t aux_gprs_count() const override; + + std::shared_ptr loop_begin_label = nullptr; + std::shared_ptr loop_end_label = nullptr; size_t work_amount = 0; + size_t wa_increment = 0; + size_t loop_id = 0; + bool evaluate_once = false; + bool is_work_amount_dynamic = false; }; -class jit_loop_begin_dynamic_emitter: public jit_loop_begin_emitter { -public: - jit_loop_begin_dynamic_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, - const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; - - void set_loop_end_label(const std::shared_ptr& label) { loop_end_label = label; } - -protected: - void validate_arguments(const std::vector &in, const std::vector &out) const override; - void emit_impl(const std::vector& in, const std::vector& out) const override; - - // For Loop arguments - size_t aux_gprs_count() const override { return 1; } - - std::shared_ptr loop_end_label; - size_t loop_id; -}; /* ============================================================== */ @@ -76,21 +54,6 @@ public: size_t get_inputs_num() const override { return 0; } -protected: - static ov::snippets::lowered::ExpressionPtr get_loop_begin_expr(const ov::snippets::lowered::ExpressionPtr& expr); - - std::shared_ptr loop_begin_label; - size_t num_inputs = 0; - size_t num_outputs = 0; - int64_t wa_increment = 0; - std::vector is_incremented = {}; -}; - -class jit_loop_end_static_emitter: public jit_loop_end_emitter { -public: - jit_loop_end_static_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; @@ -98,31 +61,25 @@ protected: void validate_arguments(const std::vector &in, const std::vector &out) const override; void emit_impl(const std::vector& in, const std::vector& out) const override; + size_t aux_gprs_count() const override; + + static ov::snippets::lowered::ExpressionPtr get_loop_begin_expr(const ov::snippets::lowered::ExpressionPtr& expr); + + std::shared_ptr loop_begin_label = nullptr; + std::shared_ptr loop_end_label = nullptr; + size_t num_inputs = 0; + size_t num_outputs = 0; size_t work_amount = 0; + size_t wa_increment = 0; std::vector is_incremented = {}; std::vector ptr_increments = {}; std::vector finalization_offsets = {}; std::vector data_sizes = {}; + size_t loop_id = 0; bool evaluate_once = false; -}; - -class jit_loop_end_dynamic_emitter: public jit_loop_end_emitter { -public: - jit_loop_end_dynamic_emitter(dnnl::impl::cpu::x64::jit_generator* h, dnnl::impl::cpu::x64::cpu_isa_t isa, - const ov::snippets::lowered::ExpressionPtr& expr); - - void emit_code(const std::vector &in_idxs, const std::vector &out_idxs, - const std::vector &pool_vec_idxs = {}, const std::vector &pool_gpr_idxs = {}) const override; - -protected: - void validate_arguments(const std::vector &in, const std::vector &out) const override; - void emit_impl(const std::vector& in, const std::vector& out) const override; - - // For Loop arguments - size_t aux_gprs_count() const override { return 1; } - - std::shared_ptr loop_end_label; - size_t loop_id; + bool are_ptr_increments_dynamic = false; + bool are_final_offsets_dynamic = false; + bool are_ptr_shifts_dynamic = false; }; /* ============================================================== */ diff --git a/src/plugins/intel_cpu/src/extension.cpp b/src/plugins/intel_cpu/src/extension.cpp index 3cf2c08cd4b..9e496ba5cd8 100644 --- a/src/plugins/intel_cpu/src/extension.cpp +++ b/src/plugins/intel_cpu/src/extension.cpp @@ -159,10 +159,8 @@ private: OP_EXTENSION(ov::snippets::op::IntermediateMemoryBuffer) \ OP_EXTENSION(ov::snippets::op::Load) \ OP_EXTENSION(ov::snippets::op::LoadReshape) \ - OP_EXTENSION(ov::snippets::op::LoopBeginStatic) \ - OP_EXTENSION(ov::snippets::op::LoopBeginDynamic) \ - OP_EXTENSION(ov::snippets::op::LoopEndStatic) \ - OP_EXTENSION(ov::snippets::op::LoopEndDynamic) \ + OP_EXTENSION(ov::snippets::op::LoopBegin) \ + OP_EXTENSION(ov::snippets::op::LoopEnd) \ OP_EXTENSION(ov::snippets::op::NewMemoryBuffer) \ OP_EXTENSION(ov::snippets::op::Nop) \ OP_EXTENSION(ov::snippets::op::PowerStatic) \