From 5933eaf11566d255e0b498e7fe739ace272b0751 Mon Sep 17 00:00:00 2001 From: intelgaoxiong Date: Tue, 11 Aug 2026 23:52:39 -0700 Subject: [PATCH] Annotate each Add in SDPA to make per-ATTN-subgraph mask-skipping decisions. Add interface to get per-SDPA's mask type. Add unit test. Signed-off-by: intelgaoxiong --- .../src/plugin/npuw/host_flash_attention.cpp | 31 ++- .../src/plugin/npuw/llm_compiled_model.cpp | 11 + .../detect_causal_mask.cpp | 62 ++++++ .../detect_causal_mask.hpp | 45 ++++ .../unit/npuw/host_flash_attention_test.cpp | 48 ++++ .../detect_causal_mask_test.cpp | 209 +++++++++++++++++- 6 files changed, 402 insertions(+), 4 deletions(-) diff --git a/src/plugins/intel_npu/src/plugin/npuw/host_flash_attention.cpp b/src/plugins/intel_npu/src/plugin/npuw/host_flash_attention.cpp index 87228370eaa18d..fc05580cc9063e 100644 --- a/src/plugins/intel_npu/src/plugin/npuw/host_flash_attention.cpp +++ b/src/plugins/intel_npu/src/plugin/npuw/host_flash_attention.cpp @@ -11,6 +11,7 @@ #include "intel_npu/ops/flash_attention_tile.hpp" #include "logging.hpp" +#include "npuw_transformations/detect_causal_mask.hpp" #include "openvino/core/validation_util.hpp" #include "openvino/op/ops.hpp" #include "openvino/openvino.hpp" @@ -1106,6 +1107,32 @@ std::optional HostFlashAttention::from(const std::shared_ptr LOG_INFO("Creating HFA tile models: tile_size=" << query_size << ", v_transposed=" << v_transposed << ", block_kv=" << block_kv_dtype << ", present_kv=" << present_kv_dtype << ", q=" << q_dtype); + + // Per-SDPA mask-skipping override + // AnnotatePerSDPAMaskType may have annotated this subgraph's Add(QK, mask) node + // with its individual mask type. For mixed SWA + global-attention models + // (e.g. Gemma-4 E2B/E4B), each ATTN subgraph decides independently: + // + // Causal: force enable_mask_skipping = true + // SlidingWindow / no annotation: keep the global enable_mask_skipping unchanged + // + // SlidingWindow is not forced to false because the global flag already handles the + // window_size >= max_prompt_len case (wide SWA that covers the full context). + bool local_enable_mask_skipping = enable_mask_skipping; + if (pattern_nodes.add_node) { + const auto& rt_info = pattern_nodes.add_node->get_rt_info(); + const auto it = rt_info.find(ov::npuw::NPUW_SDPA_MASK_TYPE_RT_KEY); + if (it != rt_info.end()) { + const auto per_sdpa_mask_type = static_cast(it->second.as()); + if (per_sdpa_mask_type == ov::npuw::MaskInfo::MaskType::Causal) { + local_enable_mask_skipping = true; + LOG_DEBUG("Per-SDPA mask annotation: Causal → mask skipping ENABLED for this ATTN subgraph"); + } else { + LOG_DEBUG("Per-SDPA mask annotation: SlidingWindow/Unknown → use global mask skipping setting (" + << (enable_mask_skipping ? "YES" : "NO") << ") for this ATTN subgraph"); + } + } + } auto tile_model = create_hfa_tile_model(q_shape_static, block_kv_dtype, // state_dtype block_kv_dtype, // kv_tile_dtype (past blocks) @@ -1115,7 +1142,7 @@ std::optional HostFlashAttention::from(const std::shared_ptr kv_num_heads, false, fused_flash_attention, - enable_mask_skipping, + local_enable_mask_skipping, v_transposed); if (!tile_model) { LOG_WARN("Failed to create HFA tile model"); @@ -1131,7 +1158,7 @@ std::optional HostFlashAttention::from(const std::shared_ptr kv_num_heads, true, fused_flash_attention, - enable_mask_skipping, + local_enable_mask_skipping, v_transposed, output_dtype); if (!final_tile_model) { diff --git a/src/plugins/intel_npu/src/plugin/npuw/llm_compiled_model.cpp b/src/plugins/intel_npu/src/plugin/npuw/llm_compiled_model.cpp index 47efb2987c72f6..8d2b0994565e4f 100644 --- a/src/plugins/intel_npu/src/plugin/npuw/llm_compiled_model.cpp +++ b/src/plugins/intel_npu/src/plugin/npuw/llm_compiled_model.cpp @@ -1051,6 +1051,17 @@ ov::npuw::LLMCompiledModel::LLMCompiledModel(const std::shared_ptr& m } else { LOG_DEBUG("Check and apply opt layout --- SKIPPED"); } + + // Annotate each Add(QK, mask) node with its per-SDPA mask type via rt_info. + // Must run AFTER OptimizeValueTensors because ScaledDotProductAttentionDecomposition + // inside it is what creates the Add(QK, mask) nodes from SDPA ops. + // For mixed SWA + global-attention models (e.g. Gemma-4 E2B/E4B), HFA uses these + // annotations to make per-ATTN-subgraph mask-skipping decisions. + ov::npuw::AnnotatePerSDPAMaskType().run_on_model(prefill_model); + for (auto& model_variant : generate_model_variants) { + ov::npuw::AnnotatePerSDPAMaskType().run_on_model(model_variant); + } + if (!m_is_embedding) { if (!m_use_chunk_prefill) { LOG_DEBUG("Removing EmptyKVInputs"); diff --git a/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.cpp b/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.cpp index abc15860d43a3b..b19116590b741a 100644 --- a/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.cpp +++ b/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.cpp @@ -5,7 +5,10 @@ #include "detect_causal_mask.hpp" #include +#include +#include +#include "../util.hpp" #include "openvino/op/ops.hpp" #include "openvino/op/scaled_dot_product_attention.hpp" #include "openvino/pass/graph_rewrite.hpp" @@ -244,4 +247,63 @@ bool DetectAttentionMask::run_on_model(const std::shared_ptr& model) return false; } +// Traces the mask input of Add(QK, mask) backward using BFS. +// Returns SlidingWindow if a Greater node appears directly inside a +// BitwiseAnd / BitwiseOr / LogicalAnd (window-size check), Causal otherwise. +static MaskInfo::MaskType detect_sdpa_mask_type_from_add(const std::shared_ptr& add_node) { + if (!add_node || add_node->get_input_size() < 2) { + return MaskInfo::MaskType::Unknown; + } + + // Trace mask input (input 1 of Add) backward using BFS. + // Look for SWA indicators: a BitwiseAnd / BitwiseOr / LogicalAnd that has + // a Greater node as a direct input (window-size check). + std::unordered_set visited; + std::queue> queue; + queue.push(add_node->get_input_node_shared_ptr(1)); + + while (!queue.empty()) { + auto node = queue.front(); + queue.pop(); + if (!node || !visited.insert(node.get()).second) + continue; + + // SWA anchor: BitwiseAnd / BitwiseOr / LogicalAnd whose direct input is Greater + if (ov::is_type(node) || ov::is_type(node) || + ov::is_type(node)) { + for (size_t i = 0; i < node->get_input_size(); ++i) { + if (ov::is_type(node->get_input_node_shared_ptr(i))) { + return MaskInfo::MaskType::SlidingWindow; + } + } + } + + // Don't cross Parameters or Constants – they are leaf nodes. + if (ov::op::util::is_parameter(node) || ov::op::util::is_constant(node)) + continue; + + for (size_t i = 0; i < node->get_input_size(); ++i) + queue.push(node->get_input_node_shared_ptr(i)); + } + + // No SWA pattern found, treat as causal. + return MaskInfo::MaskType::Causal; +} + +bool AnnotatePerSDPAMaskType::run_on_model(const std::shared_ptr& model) { + m_annotations.clear(); + + const auto all_patterns = ov::npuw::util::find_all_sdpa_pattern_nodes(model); + m_annotations.reserve(all_patterns.size()); + + for (const auto& pattern : all_patterns) { + if (!pattern.add_node) + continue; + const auto mask_type = detect_sdpa_mask_type_from_add(pattern.add_node); + pattern.add_node->get_rt_info()[NPUW_SDPA_MASK_TYPE_RT_KEY] = static_cast(mask_type); + m_annotations.push_back({pattern.add_node->get_friendly_name(), mask_type}); + } + return false; +} + } // namespace ov::npuw diff --git a/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.hpp b/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.hpp index 283de4ad216059..3e1f653b6f19b1 100644 --- a/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.hpp +++ b/src/plugins/intel_npu/src/plugin/npuw/npuw_transformations/detect_causal_mask.hpp @@ -4,6 +4,9 @@ #pragma once +#include +#include + #include "openvino/pass/pass.hpp" namespace ov::npuw { @@ -39,4 +42,46 @@ class DetectAttentionMask : public ov::pass::ModelPass { MaskInfo m_mask_info; }; +// rt_info key written by AnnotatePerSDPAMaskType and read by HostFlashAttention::from(). +// Value type: int, corresponding to MaskInfo::MaskType. +static constexpr const char* NPUW_SDPA_MASK_TYPE_RT_KEY = "npuw_sdpa_mask_type"; + +// Pre-partitioning pass: annotates each decomposed-SDPA's Add(QK, mask) node in the +// model with its individual mask type via rt_info[NPUW_SDPA_MASK_TYPE_RT_KEY]. +// +// This enables per-layer mask-skipping decisions inside HostFlashAttention::from() +// for mixed SWA + global-attention models (e.g. Gemma-4 E2B/E4B): global-attention +// ATTN subgraphs can keep mask skipping enabled even when SWA layers are present. +// +// Must be run on the whole model BEFORE partitioning so the annotation is carried +// into the isolated ATTN subgraphs (the Add node object is shared, not cloned). +// Never modifies the graph structure; run_on_model always returns false. +class AnnotatePerSDPAMaskType : public ov::pass::ModelPass { +public: + struct Annotation { + std::string add_node_name; + MaskInfo::MaskType mask_type = MaskInfo::MaskType::Unknown; + }; + + OPENVINO_MODEL_PASS_RTTI("ov::npuw::AnnotatePerSDPAMaskType"); + bool run_on_model(const std::shared_ptr& model) override; + + // Collected per-SDPA mask types from the most recent run_on_model() call. + const std::vector& get_annotations() const { + return m_annotations; + } + + // Convenience helper: returns only mask types in traversal order. + std::vector get_mask_types() const { + std::vector mask_types; + mask_types.reserve(m_annotations.size()); + for (const auto& annotation : m_annotations) + mask_types.push_back(annotation.mask_type); + return mask_types; + } + +private: + std::vector m_annotations; +}; + } // namespace ov::npuw diff --git a/src/plugins/intel_npu/tests/unit/npuw/host_flash_attention_test.cpp b/src/plugins/intel_npu/tests/unit/npuw/host_flash_attention_test.cpp index f005b12fe7f5d3..dba92c2eeb142e 100644 --- a/src/plugins/intel_npu/tests/unit/npuw/host_flash_attention_test.cpp +++ b/src/plugins/intel_npu/tests/unit/npuw/host_flash_attention_test.cpp @@ -9,6 +9,7 @@ #include #include +#include "npuw_transformations/detect_causal_mask.hpp" #include "openvino/op/add.hpp" #include "openvino/op/concat.hpp" #include "openvino/op/convert.hpp" @@ -310,6 +311,53 @@ TEST(HostFlashAttentionFromTest, Fused_MaskTileAtIndexSixInRegularTileWhenMaskSk expect_input_name(result->_final_tile_model, 6, "MASK_TILE", "fused final tile"); } +TEST(HostFlashAttentionFromTest, Fused_PerSDPACausalRtInfo_OverridesGlobalNoAndSkipsRegularMask) { + auto model = build_sdpa_model(); + ASSERT_NE(model, nullptr); + + std::shared_ptr add; + for (const auto& node : model->get_ops()) { + add = ov::as_type_ptr(node); + if (add && add->get_friendly_name() == "add.0") + break; + } + ASSERT_NE(add, nullptr); + + add->get_rt_info()[ov::npuw::NPUW_SDPA_MASK_TYPE_RT_KEY] = static_cast(ov::npuw::MaskInfo::MaskType::Causal); + + // Emulate mixed-model global decision: global NO, but this ATTN subgraph is global/causal. + auto result = ov::npuw::function::HostFlashAttention::from(model, true, false); + ASSERT_TRUE(result.has_value()); + + // Regular tile skips mask (6 inputs), final tile still keeps mask (7 inputs). + EXPECT_EQ(result->_tile_model->inputs().size(), 6u); + EXPECT_EQ(result->_final_tile_model->inputs().size(), 7u); +} + +TEST(HostFlashAttentionFromTest, Fused_PerSDPASlidingRtInfo_KeepsMaskWhenGlobalNo) { + auto model = build_sdpa_model(); + ASSERT_NE(model, nullptr); + + std::shared_ptr add; + for (const auto& node : model->get_ops()) { + add = ov::as_type_ptr(node); + if (add && add->get_friendly_name() == "add.0") + break; + } + ASSERT_NE(add, nullptr); + + add->get_rt_info()[ov::npuw::NPUW_SDPA_MASK_TYPE_RT_KEY] = + static_cast(ov::npuw::MaskInfo::MaskType::SlidingWindow); + + // Emulate mixed-model global decision: global NO and this ATTN subgraph is SWA. + auto result = ov::npuw::function::HostFlashAttention::from(model, true, false); + ASSERT_TRUE(result.has_value()); + + // Regular tile keeps mask (7 inputs), final tile always keeps mask (7 inputs). + EXPECT_EQ(result->_tile_model->inputs().size(), 7u); + EXPECT_EQ(result->_final_tile_model->inputs().size(), 7u); +} + // ============================================================================ // Tile param index map // ============================================================================ diff --git a/src/plugins/intel_npu/tests/unit/npuw/pipeline_passes/detect_causal_mask_test.cpp b/src/plugins/intel_npu/tests/unit/npuw/pipeline_passes/detect_causal_mask_test.cpp index b1008e353742fa..7719d67cffde57 100644 --- a/src/plugins/intel_npu/tests/unit/npuw/pipeline_passes/detect_causal_mask_test.cpp +++ b/src/plugins/intel_npu/tests/unit/npuw/pipeline_passes/detect_causal_mask_test.cpp @@ -7,6 +7,7 @@ #include #include +#include #include "../llm_test_helpers.hpp" #include "model_builder.hpp" @@ -14,6 +15,7 @@ #include "openvino/op/scaled_dot_product_attention.hpp" #include "openvino/openvino.hpp" +using ov::npuw::AnnotatePerSDPAMaskType; using ov::npuw::DetectAttentionMask; using ov::npuw::MaskInfo; using MaskType = ov::npuw::MaskInfo::MaskType; @@ -32,6 +34,105 @@ MaskType detect(const std::shared_ptr& model) { return pass.get_mask_info().mask_type; } +std::shared_ptr append_decomposed_sdpa_branch(int layer_idx, + bool is_sliding_mask, + ov::ParameterVector& params, + ov::ResultVector& results) { + using namespace ov::op; + + const std::string idx = std::to_string(layer_idx); + + auto q = std::make_shared(ov::element::f32, ov::Shape{1, 1, 4, 8}); + auto past_k = std::make_shared(ov::element::f32, ov::Shape{1, 1, 4, 8}); + auto present_k = std::make_shared(ov::element::f32, ov::Shape{1, 1, 4, 8}); + auto past_v = std::make_shared(ov::element::f32, ov::Shape{1, 1, 4, 8}); + auto present_v = std::make_shared(ov::element::f32, ov::Shape{1, 1, 4, 8}); + + q->set_friendly_name("query." + idx); + past_k->set_friendly_name("past_key_values." + idx + ".key"); + present_k->set_friendly_name("present." + idx + ".key"); + past_v->set_friendly_name("past_key_values." + idx + ".value"); + present_v->set_friendly_name("present." + idx + ".value"); + + params.insert(params.end(), {q, past_k, present_k, past_v, present_v}); + + auto key_concat = std::make_shared(ov::OutputVector{past_k, present_k}, 2); + key_concat->set_friendly_name("concat_key." + idx); + auto value_concat = std::make_shared(ov::OutputVector{past_v, present_v}, 2); + value_concat->set_friendly_name("concat_value." + idx); + + auto zero_i64 = v0::Constant::create(ov::element::i64, ov::Shape{}, {0}); + auto one_i64 = v0::Constant::create(ov::element::i64, ov::Shape{}, {1}); + auto four_i64 = v0::Constant::create(ov::element::i64, ov::Shape{}, {4}); // Q seq length + auto eight_i64 = + v0::Constant::create(ov::element::i64, ov::Shape{}, {8}); // KV context length (past=4 + present=4) + + // k_range covers all KV positions; q_range covers query positions only. + // k_unsq → [1, kv_len], q_unsq → [seq, 1]; comparison broadcasts to [seq, kv_len]. + auto k_range = std::make_shared(zero_i64, eight_i64, one_i64, ov::element::i64); + auto q_range = std::make_shared(zero_i64, four_i64, one_i64, ov::element::i64); + auto k_unsq = std::make_shared(k_range, zero_i64); // [1, 8] + auto q_unsq = std::make_shared(q_range, one_i64); // [4, 1] + + std::shared_ptr mask_bool; + if (is_sliding_mask) { + auto neg_window = v0::Constant::create(ov::element::i64, ov::Shape{}, {-2}); + auto bound = std::make_shared(q_unsq, neg_window); + auto greater = std::make_shared(k_unsq, bound); + auto causal = std::make_shared(k_unsq, q_unsq); + auto one_bool = v0::Constant::create(ov::element::boolean, ov::Shape{}, {true}); + auto and_win = std::make_shared(one_bool, greater); + mask_bool = std::make_shared(and_win, causal); + } else { + mask_bool = std::make_shared(k_unsq, q_unsq); + } + + auto zero_f = v0::Constant::create(ov::element::f32, ov::Shape{}, {0.0f}); + auto neg_inf = v0::Constant::create(ov::element::f32, ov::Shape{}, {-std::numeric_limits::infinity()}); + auto mask_f = std::make_shared(mask_bool, zero_f, neg_inf); + auto m0 = std::make_shared(mask_f, zero_i64); + auto mask_4d = std::make_shared(m0, zero_i64); + + auto qk = std::make_shared(q, key_concat, false, true); + qk->set_friendly_name("matmul1." + idx); + auto add = std::make_shared(qk, mask_4d); + add->set_friendly_name("add." + idx); + auto softmax = std::make_shared(add, -1); + softmax->set_friendly_name("softmax." + idx); + auto out = std::make_shared(softmax, value_concat, false, false); + out->set_friendly_name("matmul2." + idx); + + auto result = std::make_shared(out); + result->set_friendly_name("attn_out." + idx); + results.push_back(result); + return add; +} + +std::pair, std::shared_ptr> make_decomposed_sdpa_with_mask(bool sliding_mask) { + ov::ParameterVector params; + ov::ResultVector results; + auto add = append_decomposed_sdpa_branch(/*layer_idx=*/0, sliding_mask, params, results); + auto model = std::make_shared(results, params); + return {model, add}; +} + +struct MixedAttentionModel { + std::shared_ptr model; + std::shared_ptr add_global; + std::shared_ptr add_swa; +}; + +MixedAttentionModel make_mixed_decomposed_sdpa_model() { + ov::ParameterVector params; + ov::ResultVector results; + + auto add_global = append_decomposed_sdpa_branch(/*layer_idx=*/0, /*is_sliding_mask=*/false, params, results); + auto add_swa = append_decomposed_sdpa_branch(/*layer_idx=*/1, /*is_sliding_mask=*/true, params, results); + + auto model = std::make_shared(results, params); + return {model, add_global, add_swa}; +} + // --------------------------------------------------------------------------- // Minimal Q/K/V SDPA model (no explicit mask). is_causal drives the op attribute. // --------------------------------------------------------------------------- @@ -384,6 +485,110 @@ TEST(DetectAttentionMaskTest, RealPattern_Phi3Sliding_IsSlidingWindowWithSize) { // Full attention — SDPA with no mask and is_causal=false. TEST(DetectAttentionMaskTest, FullAttentionSDPA_IsUnknown) { - EXPECT_EQ(detect(make_sdpa_model(/*is_causal=*/false)), - MaskInfo::MaskType::Unknown); + EXPECT_EQ(detect(make_sdpa_model(/*is_causal=*/false)), MaskInfo::MaskType::Unknown); +} + +// ============================================================================ +// Per-SDPA annotation pass — must write rt_info and expose detected mask types +// ============================================================================ + +TEST(DetectAttentionMaskTest, AnnotatePerSDPAMaskType_Causal_WritesRtInfoAndGetter) { + auto [model, add] = make_decomposed_sdpa_with_mask(/*sliding_mask=*/false); + ASSERT_NE(model, nullptr); + ASSERT_NE(add, nullptr); + + AnnotatePerSDPAMaskType pass; + pass.run_on_model(model); + + const auto& annotations = pass.get_annotations(); + ASSERT_EQ(annotations.size(), 1u); + EXPECT_EQ(annotations[0].mask_type, MaskType::Causal); + + const auto mask_types = pass.get_mask_types(); + ASSERT_EQ(mask_types.size(), 1u); + EXPECT_EQ(mask_types[0], MaskType::Causal); + + const auto& rt_info = add->get_rt_info(); + const auto it = rt_info.find(ov::npuw::NPUW_SDPA_MASK_TYPE_RT_KEY); + ASSERT_NE(it, rt_info.end()); + EXPECT_EQ(static_cast(it->second.as()), MaskType::Causal); +} + +TEST(DetectAttentionMaskTest, AnnotatePerSDPAMaskType_SlidingWindow_WritesRtInfoAndGetter) { + auto [model, add] = make_decomposed_sdpa_with_mask(/*sliding_mask=*/true); + ASSERT_NE(model, nullptr); + ASSERT_NE(add, nullptr); + + AnnotatePerSDPAMaskType pass; + pass.run_on_model(model); + + const auto& annotations = pass.get_annotations(); + ASSERT_EQ(annotations.size(), 1u); + EXPECT_EQ(annotations[0].mask_type, MaskType::SlidingWindow); + + const auto mask_types = pass.get_mask_types(); + ASSERT_EQ(mask_types.size(), 1u); + EXPECT_EQ(mask_types[0], MaskType::SlidingWindow); + + const auto& rt_info = add->get_rt_info(); + const auto it = rt_info.find(ov::npuw::NPUW_SDPA_MASK_TYPE_RT_KEY); + ASSERT_NE(it, rt_info.end()); + EXPECT_EQ(static_cast(it->second.as()), MaskType::SlidingWindow); +} + +TEST(DetectAttentionMaskTest, AnnotatePerSDPAMaskType_ClearsResultsBetweenRuns) { + auto [model_with_pattern, _] = make_decomposed_sdpa_with_mask(/*sliding_mask=*/true); + ASSERT_NE(model_with_pattern, nullptr); + + AnnotatePerSDPAMaskType pass; + pass.run_on_model(model_with_pattern); + ASSERT_EQ(pass.get_annotations().size(), 1u); + + auto model_without_pattern = make_sdpa_model(/*is_causal=*/false); + ASSERT_NE(model_without_pattern, nullptr); + pass.run_on_model(model_without_pattern); + EXPECT_TRUE(pass.get_annotations().empty()); + EXPECT_TRUE(pass.get_mask_types().empty()); +} + +TEST(DetectAttentionMaskTest, AnnotatePerSDPAMaskType_MixedSWAAndGlobal_ReturnsBothMaskTypes) { + auto mixed = make_mixed_decomposed_sdpa_model(); + ASSERT_NE(mixed.model, nullptr); + ASSERT_NE(mixed.add_global, nullptr); + ASSERT_NE(mixed.add_swa, nullptr); + + AnnotatePerSDPAMaskType pass; + pass.run_on_model(mixed.model); + + const auto mask_types = pass.get_mask_types(); + ASSERT_EQ(mask_types.size(), 2u); + std::multiset type_set(mask_types.begin(), mask_types.end()); + EXPECT_EQ(type_set.count(MaskType::Causal), 1u); + EXPECT_EQ(type_set.count(MaskType::SlidingWindow), 1u); + + const auto& annotations = pass.get_annotations(); + ASSERT_EQ(annotations.size(), 2u); + std::multiset annotation_types; + for (const auto& annotation : annotations) + annotation_types.insert(annotation.mask_type); + EXPECT_EQ(annotation_types.count(MaskType::Causal), 1u); + EXPECT_EQ(annotation_types.count(MaskType::SlidingWindow), 1u); +} + +TEST(DetectAttentionMaskTest, AnnotatePerSDPAMaskType_MixedSWAAndGlobal_WritesCorrectRtInfo) { + auto mixed = make_mixed_decomposed_sdpa_model(); + ASSERT_NE(mixed.model, nullptr); + + AnnotatePerSDPAMaskType pass; + pass.run_on_model(mixed.model); + + auto get_mask_type_from_rt_info = [](const std::shared_ptr& add_node) { + const auto& rt_info = add_node->get_rt_info(); + const auto it = rt_info.find(ov::npuw::NPUW_SDPA_MASK_TYPE_RT_KEY); + EXPECT_NE(it, rt_info.end()); + return static_cast(it->second.as()); + }; + + EXPECT_EQ(get_mask_type_from_rt_info(mixed.add_global), MaskType::Causal); + EXPECT_EQ(get_mask_type_from_rt_info(mixed.add_swa), MaskType::SlidingWindow); }