已合并
fix(add_layer_norm): 修复InferShape未将norm多尾轴置1的bug #8804
rk创建于 8月17日
fix(add_layer_norm): 修复InferShape未将norm多尾轴置1的bug #8804
已合并
共 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 | ||
| 97 | IMPL_OP_INFERSHAPE(AddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm); | 99 | IMPL_OP_INFERSHAPE(AddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm); |
| 98 | IMPL_OP_INFERSHAPE(InplaceAddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm); | 100 | IMPL_OP_INFERSHAPE(InplaceAddLayerNorm).InferShape(InferShape4AddLayerNorm).InferDataType(InferDataType4AddLayerNorm); |
| 99 | -} // namespace ops | 101 | +} // 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 | + | ||
| 48 | TEST_F(AddLayerNorm, AddLayerNorm_InferDtype_mix_case_0) | 72 | TEST_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); |