已合并
fix(add_layer_norm): 修复InferShape未将norm多尾轴置1的bug #8804
fix(add_layer_norm): 修复InferShape未将norm多尾轴置1的bug #8804
已合并
rk创建于 8月17日
共 2 个文件变更+28-2
@@ -72,7 +72,9 @@ static ge::graphStatus InferShape4AddLayerNorm(gert::InferShapeContext* context)
72 return GRAPH_FAILED;72 return GRAPH_FAILED;
73 }73 }
74 auto shape(*x1_shape);74 auto shape(*x1_shape);
75- shape.SetDim(shape.GetDimNum() - gamma_shape->GetDimNum(), 1);75+ for (size_t i = shape.GetDimNum() - gamma_shape->GetDimNum(); i < shape.GetDimNum(); i++) {
76+ shape.SetDim(i, 1);
77+ }
76 *mean_shape = shape;78 *mean_shape = shape;
77 *rstd_shape = shape;79 *rstd_shape = shape;
78 return GRAPH_SUCCESS;80 return GRAPH_SUCCESS;
@@ -96,4 +98,4 @@ static graphStatus InferDataType4AddLayerNorm(gert::InferDataTypeContext* contex
96 98 
97IMPL_OP_INFERSHAPE(AddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm);99IMPL_OP_INFERSHAPE(AddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm);
98IMPL_OP_INFERSHAPE(InplaceAddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm);100IMPL_OP_INFERSHAPE(InplaceAddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm);
99-} // namespace ops101+} // namespace ops
@@ -45,6 +45,30 @@ TEST_F(AddLayerNorm, AddLayerNorm_infershape_case_0)
45 EXPECT_EQ(output_rstd_desc.GetShape().GetDims(), expected_rstd_shape);45 EXPECT_EQ(output_rstd_desc.GetShape().GetDims(), expected_rstd_shape);
46}46}
47 47 
48+TEST_F(AddLayerNorm, AddLayerNorm_infershape_multi_dim_gamma)
49+{
50+ ge::op::AddLayerNorm op;
51+ op.UpdateInputDesc("x1", create_desc({4, 1, 8, 16}, ge::DT_FLOAT16));
52+ op.UpdateInputDesc("x2", create_desc({4, 1, 8, 16}, ge::DT_FLOAT16));
53+ op.UpdateInputDesc("gamma", create_desc({8, 16}, ge::DT_FLOAT16));
54+ op.UpdateInputDesc("beta", create_desc({8, 16}, ge::DT_FLOAT16));
55+ 
56+ EXPECT_EQ(InferShapeTest(op), ge::GRAPH_SUCCESS);
57+ 
58+ auto output_y_desc = op.GetOutputDesc(0);
59+ auto output_mean_desc = op.GetOutputDesc(1);
60+ auto output_rstd_desc = op.GetOutputDesc(2);
61+ auto output_x_desc = op.GetOutputDesc(3);
62+ std::vector<int64_t> expected_y_shape = {4, 1, 8, 16};
63+ std::vector<int64_t> expected_x_shape = {4, 1, 8, 16};
64+ std::vector<int64_t> expected_mean_shape = {4, 1, 1, 1};
65+ std::vector<int64_t> expected_rstd_shape = {4, 1, 1, 1};
66+ EXPECT_EQ(output_y_desc.GetShape().GetDims(), expected_y_shape);
67+ EXPECT_EQ(output_x_desc.GetShape().GetDims(), expected_x_shape);
68+ EXPECT_EQ(output_mean_desc.GetShape().GetDims(), expected_mean_shape);
69+ EXPECT_EQ(output_rstd_desc.GetShape().GetDims(), expected_rstd_shape);
70+}
71+ 
48TEST_F(AddLayerNorm, AddLayerNorm_InferDtype_mix_case_0)72TEST_F(AddLayerNorm, AddLayerNorm_InferDtype_mix_case_0)
49{73{
50 ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("AddLayerNorm"), nullptr);74 ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("AddLayerNorm"), nullptr);