diff --git a/src/common/transformations/include/transformations/common_optimizations/rms_fusion.hpp b/src/common/transformations/include/transformations/common_optimizations/rms_fusion.hpp index f989c7ae47907c..d4dea9bd1b356e 100644 --- a/src/common/transformations/include/transformations/common_optimizations/rms_fusion.hpp +++ b/src/common/transformations/include/transformations/common_optimizations/rms_fusion.hpp @@ -30,7 +30,7 @@ namespace pass { class RMSFusion : public ov::pass::MatcherPass { public: OPENVINO_MATCHER_PASS_RTTI("RMSFusion"); - RMSFusion(bool force_tail_convert = true, bool enable_div_x = false, bool enable_without_gamma = false); + RMSFusion(bool force_tail_convert = true, bool enable_without_gamma = false); }; } // namespace pass diff --git a/src/common/transformations/src/transformations/common_optimizations/rms_fusion.cpp b/src/common/transformations/src/transformations/common_optimizations/rms_fusion.cpp index fd03b5ab2e35c1..076b0c49eb274a 100644 --- a/src/common/transformations/src/transformations/common_optimizations/rms_fusion.cpp +++ b/src/common/transformations/src/transformations/common_optimizations/rms_fusion.cpp @@ -28,7 +28,7 @@ namespace op_util = ov::op::util; namespace ov::pass { -RMSFusion::RMSFusion(bool force_tail_convert, bool enable_div_x, bool enable_without_gamma) { +RMSFusion::RMSFusion(bool force_tail_convert, bool enable_without_gamma) { // Detect RMS decomposition pattern // x * 1/Sqrt(ReduceMean(x^2,axes)+eps) * gamma auto x = pattern::any_input(); @@ -87,15 +87,9 @@ RMSFusion::RMSFusion(bool force_tail_convert, bool enable_div_x, bool enable_wit // x * 1/Sqrt(ReduceMean(x^2,axes)+eps) auto mul1 = pattern::wrap_type({x, div_or_pow}); - std::shared_ptr mul_or_div; - // TODO: Check div_x pattern failed in CPU CI Pytorch layer test. - if (enable_div_x) { - // x / Sqrt(ReduceMean(x^2,axes)+eps) - auto div_x = pattern::wrap_type({x, sqrt}); - mul_or_div = std::make_shared(OutputVector{mul1, div_x}); - } else { - mul_or_div = std::make_shared(OutputVector{mul1}); - } + // x / Sqrt(ReduceMean(x^2,axes)+eps) + auto div_x = pattern::wrap_type({x, sqrt}); + auto mul_or_div = std::make_shared(OutputVector{mul1, div_x}); // Pattern 1: RMS with gamma (learnable parameter) // x * 1/Sqrt(ReduceMean(x^2,axes)+eps) * gamma (gamma is constant) @@ -110,7 +104,9 @@ RMSFusion::RMSFusion(bool force_tail_convert, bool enable_div_x, bool enable_wit // This allows partial fusion: only fuse up to mul_or_div auto scale = pattern::any_input(pattern::class_other_than()); auto mul_with_scale = pattern::wrap_type({mul_or_div, scale}); - rms_mul = std::make_shared(OutputVector{mul_with_gamma, mul_with_scale}); + // Pattern 3: RMS without gamma (unit normalization, e.g. Gemma v_norm) + // x * 1/Sqrt(ReduceMean(x^2,axes)+eps) — no trailing Multiply + rms_mul = std::make_shared(OutputVector{mul_with_gamma, mul_with_scale, mul_or_div}); } else { rms_mul = mul_with_gamma; } @@ -139,6 +135,14 @@ RMSFusion::RMSFusion(bool force_tail_convert, bool enable_div_x, bool enable_wit auto mul_or_div_node = pattern_map.at(mul_or_div).get_node_shared_ptr(); bool elementwise_affine = pattern_map.count(mul_with_gamma); + // Avoid partial fusion (fusing without gamma) when gamma Multiply follows + if (!elementwise_affine && m.get_match_root() == mul_or_div_node) { + for (auto& target : mul_or_div_node->output(0).get_target_inputs()) { + if (ov::is_type(target.get_node())) + return false; + } + } + std::shared_ptr gamma_node; if (elementwise_affine) { gamma_node = pattern_map.at(gamma).get_node_shared_ptr(); diff --git a/src/common/transformations/tests/common_optimizations/rms_norm_decomposition_test.cpp b/src/common/transformations/tests/common_optimizations/rms_norm_decomposition_test.cpp index 8b6eba057674ab..18336ad29fb0eb 100644 --- a/src/common/transformations/tests/common_optimizations/rms_norm_decomposition_test.cpp +++ b/src/common/transformations/tests/common_optimizations/rms_norm_decomposition_test.cpp @@ -273,7 +273,7 @@ TEST_F(TransformationTestsF, RMSNormFusionTest8) { auto comp = std::make_shared(mul, ov::element::f16); model = std::make_shared(ov::OutputVector{comp}, ov::ParameterVector{input}); - manager.register_pass(true, true); + manager.register_pass(true); } { auto input = std::make_shared(ov::element::f32, ov::Shape{1, 2, 6}); @@ -307,7 +307,7 @@ TEST_F(TransformationTestsF, RMSNormFusionTest9) { auto mul = std::make_shared(gamma, div); model = std::make_shared(ov::OutputVector{mul}, ov::ParameterVector{input}); - manager.register_pass(false, true); + manager.register_pass(false); } { auto input = std::make_shared(ov::element::f32, ov::Shape{1, 2, 6}); @@ -342,7 +342,7 @@ TEST_F(TransformationTestsF, RMSNormFusionTest10) { auto mul2 = std::make_shared(mul1, scale); model = std::make_shared(ov::OutputVector{mul2}, ov::ParameterVector{input, scale}); - manager.register_pass(false, false, true); + manager.register_pass(false, true); } { auto input = std::make_shared(ov::element::f32, ov::Shape{1, 2, 6}); @@ -374,7 +374,7 @@ TEST_F(TransformationTestsF, RMSNormFusionTest11) { auto mul2 = std::make_shared(mul1, scale); model = std::make_shared(ov::OutputVector{mul2}, ov::ParameterVector{input, scale}); - manager.register_pass(false, false, true); + manager.register_pass(false, true); } { auto input = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 6}); @@ -517,7 +517,7 @@ TEST_F(TransformationTestsF, RMSNormFusionTest15_PowerNegHalf_NoGammaScale) { auto mul2 = std::make_shared(mul1, scale); model = std::make_shared(ov::OutputVector{mul2}, ov::ParameterVector{input, scale}); - manager.register_pass(false, false, true); + manager.register_pass(false, true); } { auto input = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 6}); @@ -655,7 +655,7 @@ TEST_F(TransformationTestsF, RMSNormFusionTest19_MulSquare_NoGamma_DynamicScale) auto mul2 = std::make_shared(mul1, scale); model = std::make_shared(ov::OutputVector{mul2}, ov::ParameterVector{input, scale}); - manager.register_pass(false, false, true); + manager.register_pass(false, true); } { auto input = std::make_shared(ov::element::f32, ov::PartialShape{-1, -1, 2048}); @@ -666,4 +666,29 @@ TEST_F(TransformationTestsF, RMSNormFusionTest19_MulSquare_NoGamma_DynamicScale) model_ref = std::make_shared(ov::OutputVector{mul}, ov::ParameterVector{input, scale}); } comparator.enable(FunctionsComparator::CmpValues::ACCURACY); +} + +// Power(-0.5) direct path without gamma, no trailing Multiply (unit normalization) +TEST_F(TransformationTestsF, RMSNormFusionTest20_PowerNegHalf_NoGamma) { + { + auto input = std::make_shared(ov::element::f32, ov::Shape{1, 8, 256}); + auto power_const = ov::op::v0::Constant::create(ov::element::f32, ov::Shape{}, {2.f}); + auto power = std::make_shared(input, power_const); + auto mean_axes = ov::op::v0::Constant::create(ov::element::i64, ov::Shape{1}, {-1}); + auto mean = std::make_shared(power, mean_axes, true); + auto eps = ov::op::v0::Constant::create(ov::element::f32, ov::Shape{}, {1e-6f}); + auto add_eps = std::make_shared(mean, eps); + auto neg_half = ov::op::v0::Constant::create(ov::element::f32, ov::Shape{}, {-0.5f}); + auto rsqrt = std::make_shared(add_eps, neg_half); + auto mul = std::make_shared(input, rsqrt); + + model = std::make_shared(ov::OutputVector{mul}, ov::ParameterVector{input}); + manager.register_pass(false, true); + } + { + auto input = std::make_shared(ov::element::f32, ov::Shape{1, 8, 256}); + auto rms = std::make_shared(input, 1e-6f, ov::element::f32); + model_ref = std::make_shared(ov::OutputVector{rms}, ov::ParameterVector{input}); + } + comparator.enable(FunctionsComparator::CmpValues::ACCURACY); } \ No newline at end of file diff --git a/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp b/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp index ffd8f06110dd17..bbe8decd8b0fa3 100644 --- a/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp +++ b/src/plugins/intel_gpu/src/plugin/transformations_pipeline.cpp @@ -747,7 +747,7 @@ void TransformationsPipeline::apply(std::shared_ptr func) { const int32_t vec_size = 8; return static_cast((gamma_shape.back() / vec_size)) > static_cast(device_info.max_work_group_size); }); - manager.register_pass(false, true, true); + manager.register_pass(false, true); manager.register_pass(); manager.register_pass(); manager.register_pass();