已合并
add StatelessNormal/StatelessTruncatedNormalV2 opapi/ophost UT #4434
梅国晗954517创建于 19 天前
add StatelessNormal/StatelessTruncatedNormalV2 opapi/ophost UT #4434
已合并
梅国晗954517创建于 19 天前
3 个文件变更+868-0
Arandom/stateless_normal/tests/ut/op_api/test_aclnn_stateless_normal_l0.cpp+197-0
@@ -0,0 +1,197 @@
1+/**
2+ * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file test_aclnn_stateless_normal_l0.cpp
13+ * \brief l0op::StatelessNormal L0 UT
14+ *
15+ * 该算子为 aclnn_exclude(无 L2 入口),被测对象是 op_api/stateless_normal.cpp 的两个 L0 重载:
16+ * A) 标量接口:StatelessNormal(result, seed, offset, mean, stdev, executor)
17+ * B) Tensor接口:StatelessNormal(result, seedTensor, offsetTensor, mean, stdev, executor)
18+ *
19+ * 文件名必须以 test_aclnn_ 开头 —— cmake/ut.cmake 中 OP_API_MODULE_NAME 的 glob 是
20+ * `${MODULE_DIR}/test_aclnn_*.cpp`,命名不符则不会被编入 math_op_api_ut。
21+ *
22+ * 覆盖目标(两个重载各自内部经过 StatelessNormalAiCore,输出 dtype 与 result 输入保持一致):
23+ * A-float32 : scalar_float32 A-float16 : scalar_float16 A-bfloat16 : scalar_bf16
24+ * B-float32 : tensor_float32 B-float16 : tensor_float16 B-bfloat16 : tensor_bf16
25+ *
26+ * 附加覆盖:
27+ * 输出 dtype 与输入一致(非 float32 不被静默提升): scalar_dtype_preserved
28+ * 输出 shape 与输入一致 : scalar_shape_preserved
29+ * 多维输入(ToShapeVector 多维路径) : scalar_multi_dim
30+ * 非零 seed/offset(ConvertToTensor 标量路径) : scalar_nonzero_seed_offset
31+ */
32+ 
33+#include <gtest/gtest.h>
34+#include "opdev/make_op_executor.h"
35+#include "opdev/platform.h"
36+#include "random/stateless_normal/op_api/stateless_normal.h"
37+ 
38+using namespace op;
39+using namespace std;
40+ 
41+namespace {
42+constexpr int64_t DATA_SIZE = 256;
43+constexpr int64_t DEFAULT_SEED = 0;
44+constexpr int64_t DEFAULT_OFFSET = 0;
45+} // namespace
46+ 
47+class StatelessNormalL0Test : public ::testing::Test {
48+public:
49+ StatelessNormalL0Test() : exe(nullptr) {}
50+ 
51+ aclTensor* CreateAclTensor(std::vector<int64_t> shape, aclDataType dtype)
52+ {
53+ return aclCreateTensor(shape.data(), shape.size(), dtype, nullptr, 0, ACL_FORMAT_ND, shape.data(), shape.size(),
54+ data);
55+ }
56+ 
57+ void SetUp() override
58+ {
59+ // 该算子仅注册 ascend950,统一在 950 平台下构图
60+ SetPlatformNpuArch(NpuArch::DAV_3510);
61+ auto executor = &exe;
62+ auto unique_executor = CREATE_EXECUTOR();
63+ unique_executor.ReleaseTo(executor);
64+ }
65+ 
66+ void TearDown() override { delete exe; }
67+ 
68+public:
69+ aclOpExecutor* exe;
70+ float data[DATA_SIZE] = {1.0f};
71+};
72+ 
73+// ===== 重载 A:标量 seed/offset =====
74+ 
75+// case 1: float32
76+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_float32)
77+{
78+ auto result = CreateAclTensor({256}, ACL_FLOAT);
79+ auto mean = CreateAclTensor({256}, ACL_FLOAT);
80+ auto stdev = CreateAclTensor({256}, ACL_FLOAT);
atomgit-bot
atomgit-botatomgit-bot19 天前

🟡 Medium Priority

test_aclnn_stateless_normal_l0.cpp 中,CreateAclTensor 方法内部调用 aclCreateTensor,该方法可能返回 nullptr(例如资源不足时)。但所有 10 个测试用例在调用 CreateAclTensor 后均未对返回值做空指针检查,直接将返回值传给 l0op::StatelessNormal。若任一 CreateAclTensor 调用失败返回 nullptr,后续以 nullptr 作为 tensor 参数调用 StatelessNormal 会导致崩溃或未定义行为。

影响范围:文件内所有测试用例(line 78-80, 89-91, 100-102, 113-117, 126-130, 139-143, 154-156, 168-170, 179-181, 190-192),每个 case 有 3~5 个未检查的 CreateAclTensor 调用。

建议:在每个 CreateAclTensor 调用后添加 ASSERT_NE(xxx, nullptr) 检查,确保在 tensor 创建失败时测试尽早终止并给出明确诊断,而非在后续代码中空指针解引用崩溃。

改动建议
80
+ auto result = CreateAclTensor({256}, ACL_FLOAT);
81
+ ASSERT_NE(result, nullptr);
82
+ auto mean = CreateAclTensor({256}, ACL_FLOAT);
83
+ ASSERT_NE(mean, nullptr);
80
- auto stdev = CreateAclTensor({256}, ACL_FLOAT);
84
+ auto stdev = CreateAclTensor({256}, ACL_FLOAT);
85
+ ASSERT_NE(stdev, nullptr);
应用建议
likedislike
81+ auto out = l0op::StatelessNormal(result, DEFAULT_SEED, DEFAULT_OFFSET, mean, stdev, exe);
82+ ASSERT_NE(out, nullptr);
83+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_FLOAT);
84+}
85+ 
86+// case 2: float16
87+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_float16)
88+{
89+ auto result = CreateAclTensor({256}, ACL_FLOAT16);
90+ auto mean = CreateAclTensor({256}, ACL_FLOAT16);
91+ auto stdev = CreateAclTensor({256}, ACL_FLOAT16);
92+ auto out = l0op::StatelessNormal(result, DEFAULT_SEED, DEFAULT_OFFSET, mean, stdev, exe);
93+ ASSERT_NE(out, nullptr);
94+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_FLOAT16);
95+}
96+ 
97+// case 3: bfloat16
98+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_bf16)
99+{
100+ auto result = CreateAclTensor({256}, ACL_BF16);
101+ auto mean = CreateAclTensor({256}, ACL_BF16);
102+ auto stdev = CreateAclTensor({256}, ACL_BF16);
103+ auto out = l0op::StatelessNormal(result, DEFAULT_SEED, DEFAULT_OFFSET, mean, stdev, exe);
104+ ASSERT_NE(out, nullptr);
105+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_BF16);
106+}
107+ 
108+// ===== 重载 B:Tensor 形式 seed/offset(图捕获模式)=====
109+ 
110+// case 4: float32
111+TEST_F(StatelessNormalL0Test, stateless_normal_l0_tensor_float32)
112+{
113+ auto result = CreateAclTensor({256}, ACL_FLOAT);
114+ auto seedTensor = CreateAclTensor({1}, ACL_INT64);
115+ auto offsetTensor = CreateAclTensor({1}, ACL_INT64);
116+ auto mean = CreateAclTensor({256}, ACL_FLOAT);
117+ auto stdev = CreateAclTensor({256}, ACL_FLOAT);
118+ auto out = l0op::StatelessNormal(result, seedTensor, offsetTensor, mean, stdev, exe);
119+ ASSERT_NE(out, nullptr);
120+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_FLOAT);
121+}
122+ 
123+// case 5: float16
124+TEST_F(StatelessNormalL0Test, stateless_normal_l0_tensor_float16)
125+{
126+ auto result = CreateAclTensor({256}, ACL_FLOAT16);
127+ auto seedTensor = CreateAclTensor({1}, ACL_INT64);
128+ auto offsetTensor = CreateAclTensor({1}, ACL_INT64);
129+ auto mean = CreateAclTensor({256}, ACL_FLOAT16);
130+ auto stdev = CreateAclTensor({256}, ACL_FLOAT16);
131+ auto out = l0op::StatelessNormal(result, seedTensor, offsetTensor, mean, stdev, exe);
132+ ASSERT_NE(out, nullptr);
133+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_FLOAT16);
134+}
135+ 
136+// case 6: bfloat16
137+TEST_F(StatelessNormalL0Test, stateless_normal_l0_tensor_bf16)
138+{
139+ auto result = CreateAclTensor({256}, ACL_BF16);
140+ auto seedTensor = CreateAclTensor({1}, ACL_INT64);
141+ auto offsetTensor = CreateAclTensor({1}, ACL_INT64);
142+ auto mean = CreateAclTensor({256}, ACL_BF16);
143+ auto stdev = CreateAclTensor({256}, ACL_BF16);
144+ auto out = l0op::StatelessNormal(result, seedTensor, offsetTensor, mean, stdev, exe);
145+ ASSERT_NE(out, nullptr);
146+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_BF16);
147+}
148+ 
149+// ===== 形状 / 参数覆盖 =====
150+ 
151+// case 7: 多维输入,ToShapeVector 多维路径,输出 shape 应与输入一致
152+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_multi_dim)
153+{
154+ auto result = CreateAclTensor({2, 8, 16}, ACL_FLOAT);
155+ auto mean = CreateAclTensor({2, 8, 16}, ACL_FLOAT);
156+ auto stdev = CreateAclTensor({2, 8, 16}, ACL_FLOAT);
157+ auto out = l0op::StatelessNormal(result, DEFAULT_SEED, DEFAULT_OFFSET, mean, stdev, exe);
158+ ASSERT_NE(out, nullptr);
159+ EXPECT_EQ(out->GetViewShape().GetDimNum(), 3U);
160+ EXPECT_EQ(out->GetViewShape().GetDim(0), 2);
161+ EXPECT_EQ(out->GetViewShape().GetDim(1), 8);
162+ EXPECT_EQ(out->GetViewShape().GetDim(2), 16);
163+}
164+ 
165+// case 8: 输出 shape 与输入严格一致(2D)
166+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_shape_preserved)
167+{
168+ auto result = CreateAclTensor({16, 16}, ACL_FLOAT);
169+ auto mean = CreateAclTensor({16, 16}, ACL_FLOAT);
170+ auto stdev = CreateAclTensor({16, 16}, ACL_FLOAT);
171+ auto out = l0op::StatelessNormal(result, DEFAULT_SEED, DEFAULT_OFFSET, mean, stdev, exe);
172+ ASSERT_NE(out, nullptr);
173+ EXPECT_EQ(out->GetViewShape(), result->GetViewShape());
174+}
175+ 
176+// case 9: 非零 seed/offset,确认标量 ConvertToTensor 路径
177+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_nonzero_seed_offset)
178+{
179+ auto result = CreateAclTensor({256}, ACL_FLOAT);
180+ auto mean = CreateAclTensor({256}, ACL_FLOAT);
181+ auto stdev = CreateAclTensor({256}, ACL_FLOAT);
182+ auto out = l0op::StatelessNormal(result, 12345, 678, mean, stdev, exe);
183+ ASSERT_NE(out, nullptr);
184+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_FLOAT);
185+}
186+ 
187+// case 10: 非 float32 输入不应被静默提升为 float32(fp16 入 → fp16 出)
188+TEST_F(StatelessNormalL0Test, stateless_normal_l0_scalar_dtype_preserved)
189+{
190+ auto result = CreateAclTensor({4, 64}, ACL_FLOAT16);
191+ auto mean = CreateAclTensor({4, 64}, ACL_FLOAT16);
192+ auto stdev = CreateAclTensor({4, 64}, ACL_FLOAT16);
193+ auto out = l0op::StatelessNormal(result, DEFAULT_SEED, DEFAULT_OFFSET, mean, stdev, exe);
194+ ASSERT_NE(out, nullptr);
195+ EXPECT_EQ(out->GetDataType(), op::DataType::DT_FLOAT16);
196+ EXPECT_EQ(out->GetViewShape(), result->GetViewShape());
197+}
Arandom/stateless_normal/tests/ut/op_host/test_stateless_normal_infershape.cpp+286-0
@@ -0,0 +1,286 @@
1+/**
2+ * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file test_stateless_normal_infershape.cpp
13+ * \brief StatelessNormal InferShape UT
14+ *
15+ * 被测实现:op_host/stateless_normal_infershape.cpp
16+ * → ops::randomCommon::CommonInferShape(context, {{"shape",0}}, {{"y",0}}, MODE_DEPENDENCY)
17+ *
18+ * 关键语义(random_infershape_base.cpp):
19+ * xShapeSize = inShape->GetShapeSize() ← 输入 shape tensor 的元素个数 = 输出 rank
20+ * DependencyMode() 按 shape tensor 的 dtype 分派:
21+ * DT_INT64 → HandleShapeTensor<int64_t>,逐维写入 const
22+ * DT_INT32 → HandleShapeTensor<int32_t>
23+ * 其它 → return false → GRAPH_FAILED
24+ * const data 为 nullptr 时 → SetUnknownShape(xShapeSize) → 各维 -1
25+ *
26+ * 输入布局(与 op_def / tiling UT 一致):
27+ * [0] shape: DT_INT64, 1D, const(值依赖,InputsDataDependency({0}))
28+ * [1] seed: DT_INT64 scalar
29+ * [2] offset: DT_INT64 scalar
30+ * [3] mean: DT_FLOAT/DT_FLOAT16/DT_BF16 tensor(与输出同 shape)
31+ * [4] stdev: DT_FLOAT/DT_FLOAT16/DT_BF16 tensor(与输出同 shape)
32+ * 属性:dtype(OPTIONAL Int,0=float32 / 2=float16 / 3=bfloat16)
33+ *
34+ * 覆盖矩阵:
35+ * 1D 输出 : case_1d
36+ * 2D 输出 : case_2d
37+ * 4D 输出 : case_4d
38+ * 含 1 的维度 : case_dim_with_one
39+ * float16 / bfloat16 输出 dtype: case_fp16 / case_bf16
40+ * INT32 shape tensor 分支 : case_int32_shape
41+ * const data 缺失 → unknown : case_null_const_1d / case_null_const_2d
42+ * 非法 shape dtype → FAILED : case_invalid_shape_dtype
43+ */
44+ 
45+#include <gtest/gtest.h>
46+#include <iostream>
47+#include "infershape_context_faker.h"
48+#include "infershape_case_executor.h"
49+ 
50+using namespace std;
51+ 
52+class StatelessNormalInferShapeTest : public testing::Test {
53+protected:
54+ static void SetUpTestCase() { std::cout << "StatelessNormalInferShapeTest SetUp" << std::endl; }
55+ 
56+ static void TearDownTestCase() { std::cout << "StatelessNormalInferShapeTest TearDown" << std::endl; }
57+};
58+ 
59+// case 1: 1D 输出。shape tensor = [16384],1 个元素 → 输出 rank 1
60+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_1d)
61+{
62+ vector<int64_t> shapeValue = {16384};
63+ gert::InfershapeContextPara infershapeContextPara(
64+ "StatelessNormal",
65+ {
66+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
67+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
68+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
69+ {{{16384}, {16384}}, ge::DT_FLOAT, ge::FORMAT_ND},
70+ {{{16384}, {16384}}, ge::DT_FLOAT, ge::FORMAT_ND},
71+ },
72+ {
73+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
74+ },
75+ {
76+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
77+ });
78+ std::vector<std::vector<int64_t>> expectOutputShape = {{16384}};
79+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
80+}
81+ 
82+// case 2: 2D 输出。shape tensor 含 2 个元素 → 输出 rank 2
83+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_2d)
84+{
85+ vector<int64_t> shapeValue = {32, 512};
86+ gert::InfershapeContextPara infershapeContextPara(
87+ "StatelessNormal",
88+ {
89+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
90+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
91+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
92+ {{{16384}, {16384}}, ge::DT_FLOAT, ge::FORMAT_ND},
93+ {{{16384}, {16384}}, ge::DT_FLOAT, ge::FORMAT_ND},
94+ },
95+ {
96+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
97+ },
98+ {
99+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
100+ });
101+ std::vector<std::vector<int64_t>> expectOutputShape = {{32, 512}};
102+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
103+}
104+ 
105+// case 3: 4D 输出,验证高维逐维写入正确
106+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_4d)
107+{
108+ vector<int64_t> shapeValue = {2, 8, 32, 32};
109+ gert::InfershapeContextPara infershapeContextPara(
110+ "StatelessNormal",
111+ {
112+ {{{4}, {4}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
113+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
114+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
115+ {{{16384}, {16384}}, ge::DT_FLOAT, ge::FORMAT_ND},
116+ {{{16384}, {16384}}, ge::DT_FLOAT, ge::FORMAT_ND},
117+ },
118+ {
119+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
120+ },
121+ {
122+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
123+ });
124+ std::vector<std::vector<int64_t>> expectOutputShape = {{2, 8, 32, 32}};
125+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
126+}
127+ 
128+// case 4: 含 1 的维度,确认不会被误当作标量或被压缩掉
129+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_dim_with_one)
130+{
131+ vector<int64_t> shapeValue = {1, 4097, 1};
132+ gert::InfershapeContextPara infershapeContextPara(
133+ "StatelessNormal",
134+ {
135+ {{{3}, {3}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
136+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
137+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
138+ {{{4097}, {4097}}, ge::DT_FLOAT, ge::FORMAT_ND},
139+ {{{4097}, {4097}}, ge::DT_FLOAT, ge::FORMAT_ND},
140+ },
141+ {
142+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
143+ },
144+ {
145+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
146+ });
147+ std::vector<std::vector<int64_t>> expectOutputShape = {{1, 4097, 1}};
148+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
149+}
150+ 
151+// case 5: float16 输出(dtype attr = 2)。infershape 不依赖 dtype,但保证该组合能通
152+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_fp16)
153+{
154+ vector<int64_t> shapeValue = {64, 256};
155+ gert::InfershapeContextPara infershapeContextPara(
156+ "StatelessNormal",
157+ {
158+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
159+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
160+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
161+ {{{16384}, {16384}}, ge::DT_FLOAT16, ge::FORMAT_ND},
162+ {{{16384}, {16384}}, ge::DT_FLOAT16, ge::FORMAT_ND},
163+ },
164+ {
165+ {{{}, {}}, ge::DT_FLOAT16, ge::FORMAT_ND},
166+ },
167+ {
168+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(2)},
169+ });
170+ std::vector<std::vector<int64_t>> expectOutputShape = {{64, 256}};
171+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
172+}
173+ 
174+// case 6: bfloat16 输出(dtype attr = 3)
175+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_bf16)
176+{
177+ vector<int64_t> shapeValue = {32, 512};
178+ gert::InfershapeContextPara infershapeContextPara(
179+ "StatelessNormal",
180+ {
181+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
182+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
183+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
184+ {{{16384}, {16384}}, ge::DT_BF16, ge::FORMAT_ND},
185+ {{{16384}, {16384}}, ge::DT_BF16, ge::FORMAT_ND},
186+ },
187+ {
188+ {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND},
189+ },
190+ {
191+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(3)},
192+ });
193+ std::vector<std::vector<int64_t>> expectOutputShape = {{32, 512}};
194+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
195+}
196+ 
197+// case 7: INT32 shape tensor → DependencyMode 的 HandleShapeTensor<int32_t> 分支
198+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_int32_shape)
199+{
200+ vector<int32_t> shapeValue = {16, 64};
201+ gert::InfershapeContextPara infershapeContextPara(
202+ "StatelessNormal",
203+ {
204+ {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true, shapeValue.data()},
205+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
206+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
207+ {{{1024}, {1024}}, ge::DT_FLOAT, ge::FORMAT_ND},
208+ {{{1024}, {1024}}, ge::DT_FLOAT, ge::FORMAT_ND},
209+ },
210+ {
211+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
212+ },
213+ {
214+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
215+ });
216+ std::vector<std::vector<int64_t>> expectOutputShape = {{16, 64}};
217+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
218+}
219+ 
220+// case 8: const data 缺失(1 元素)→ GetData 返回 nullptr → SetUnknownShape(1) → {-1}
221+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_null_const_1d)
222+{
223+ vector<int64_t> shapeValue = {};
224+ gert::InfershapeContextPara infershapeContextPara(
225+ "StatelessNormal",
226+ {
227+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
228+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
229+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
230+ {{{1}, {1}}, ge::DT_FLOAT, ge::FORMAT_ND},
231+ {{{1}, {1}}, ge::DT_FLOAT, ge::FORMAT_ND},
232+ },
233+ {
234+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
235+ },
236+ {
237+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
238+ });
239+ std::vector<std::vector<int64_t>> expectOutputShape = {{-1}};
240+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
241+}
242+ 
243+// case 9: const data 缺失,shape tensor 有 2 个元素 → SetUnknownShape(2) → {-1,-1}
244+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_null_const_2d)
245+{
246+ vector<int64_t> shapeValue = {};
247+ gert::InfershapeContextPara infershapeContextPara(
248+ "StatelessNormal",
249+ {
250+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
251+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
252+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
253+ {{{2}, {2}}, ge::DT_FLOAT, ge::FORMAT_ND},
254+ {{{2}, {2}}, ge::DT_FLOAT, ge::FORMAT_ND},
255+ },
256+ {
257+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
258+ },
259+ {
260+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
261+ });
262+ std::vector<std::vector<int64_t>> expectOutputShape = {{-1, -1}};
263+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
264+}
265+ 
266+// case 10: 非法 shape dtype(DT_FLOAT)→ DependencyMode 两个分支都不命中 → GRAPH_FAILED
267+TEST_F(StatelessNormalInferShapeTest, stateless_normal_infershape_case_invalid_shape_dtype)
268+{
269+ vector<float> shapeValue = {32.0f, 512.0f};
270+ gert::InfershapeContextPara infershapeContextPara(
271+ "StatelessNormal",
272+ {
273+ {{{2}, {2}}, ge::DT_FLOAT, ge::FORMAT_ND, true, shapeValue.data()},
274+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
275+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND},
276+ {{{1024}, {1024}}, ge::DT_FLOAT, ge::FORMAT_ND},
277+ {{{1024}, {1024}}, ge::DT_FLOAT, ge::FORMAT_ND},
278+ },
279+ {
280+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
281+ },
282+ {
283+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
284+ });
285+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED);
286+}
Arandom/stateless_truncated_normal_v2/tests/ut/op_host/test_stateless_truncated_normal_v2_infershape.cpp+385-0
@@ -0,0 +1,385 @@
1+/**
2+ * Copyright (c) 2025-2026 Huawei Technologies Co., Ltd.
3+ * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4+ * CANN Open Software License Agreement Version 2.0 (the "License").
5+ * Please refer to the License for details. You may not use this file except in compliance with the License.
6+ * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7+ * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8+ * See LICENSE in the root of the software repository for the full text of the License.
9+ */
10+ 
11+/*!
12+ * \file test_stateless_truncated_normal_v2_infershape.cpp
13+ * \brief StatelessTruncatedNormalV2 InferShape UT
14+ *
15+ * 被测实现:op_host/stateless_truncated_normal_v2_infershape.cpp
16+ * → ops::randomCommon::CommonInferShape(context, {{"shape",0}}, {{"y",0}}, MODE_DEPENDENCY)
17+ *
18+ * 关键语义(random_infershape_base.cpp):
19+ * xShapeSize = inShape->GetShapeSize() ← 输入 shape tensor 的元素个数 = 输出 rank
20+ * DependencyMode() 按 shape tensor 的 dtype 分派:
21+ * DT_INT64 → HandleShapeTensor<int64_t>,逐维写入 const
22+ * DT_INT32 → HandleShapeTensor<int32_t>
23+ * 其它 → return false → GRAPH_FAILED
24+ * const data 为 nullptr 时 → SetUnknownShape(xShapeSize) → 各维 -1
25+ *
26+ * 输入布局(与 op_def / tiling UT 一致):
27+ * [0] shape: DT_INT32/DT_INT64, 1D, const(值依赖,InputsDataDependency({0}))
28+ * [1] key: DT_UINT64, shape {1}
29+ * [2] counter: DT_UINT64, shape {2}
30+ * [3] alg: DT_INT32 scalar(ALG_PHILOX = 1
31+ * 属性:dtype(OPTIONAL Int,0=float32 / 1=float16 / 27=bfloat16,与本算子 tiling UT 一致的取值约定)
32+ *
33+ * infershape 逻辑仅依赖 shape(index 0),key/counter/alg 的取值不影响推导结果,
34+ * 这里填入与 tiling UT 一致的合法哑值,仅用于保持输入布局完整。
35+ *
36+ * 覆盖矩阵:
37+ * 1D 输出 : case_1d
38+ * 2D 输出 : case_2d
39+ * 4D 输出 : case_4d
40+ * 5D 输出 : case_5d
41+ * 含 1 的维度 : case_dim_with_one
42+ * float16 / bfloat16 输出 dtype: case_fp16 / case_bf16
43+ * INT32 / INT64 shape tensor : case_int32_shape / case_int64_shape
44+ * const data 缺失 → unknown : case_null_const_1d / case_null_const_2d
45+ * 0 元素 shape tensor(标量输出): case_0dim_scalar
46+ * 非法 shape dtype → FAILED : case_invalid_shape_dtype
47+ */
48+ 
49+#include <gtest/gtest.h>
50+#include <iostream>
51+#include "infershape_context_faker.h"
52+#include "infershape_case_executor.h"
53+ 
54+using namespace std;
55+ 
56+class StatelessTruncatedNormalV2InferShapeTest : public testing::Test {
57+protected:
58+ static void SetUpTestCase() { std::cout << "StatelessTruncatedNormalV2InferShapeTest SetUp" << std::endl; }
59+ 
60+ static void TearDownTestCase() { std::cout << "StatelessTruncatedNormalV2InferShapeTest TearDown" << std::endl; }
61+};
62+ 
63+// case 1: 1D 输出。shape tensor = [16384],1 个元素 → 输出 rank 1
64+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_1d)
65+{
66+ vector<int64_t> shapeValue = {16384};
67+ uint64_t keyValue = 42;
68+ uint64_t counterValue[2] = {0, 0};
69+ int32_t algValue = 1;
70+ gert::InfershapeContextPara infershapeContextPara(
71+ "StatelessTruncatedNormalV2",
72+ {
73+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
74+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
75+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
76+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
77+ },
78+ {
79+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
80+ },
81+ {
82+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
83+ });
84+ std::vector<std::vector<int64_t>> expectOutputShape = {{16384}};
85+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
86+}
87+ 
88+// case 2: 2D 输出。shape tensor 含 2 个元素 → 输出 rank 2
89+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_2d)
90+{
91+ vector<int64_t> shapeValue = {32, 512};
92+ uint64_t keyValue = 42;
93+ uint64_t counterValue[2] = {0, 0};
94+ int32_t algValue = 1;
95+ gert::InfershapeContextPara infershapeContextPara(
96+ "StatelessTruncatedNormalV2",
97+ {
98+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
99+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
100+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
101+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
102+ },
103+ {
104+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
105+ },
106+ {
107+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
108+ });
109+ std::vector<std::vector<int64_t>> expectOutputShape = {{32, 512}};
110+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
111+}
112+ 
113+// case 3: 4D 输出,验证高维逐维写入正确
114+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_4d)
115+{
116+ vector<int64_t> shapeValue = {2, 8, 32, 32};
117+ uint64_t keyValue = 42;
118+ uint64_t counterValue[2] = {0, 0};
119+ int32_t algValue = 1;
120+ gert::InfershapeContextPara infershapeContextPara(
121+ "StatelessTruncatedNormalV2",
122+ {
123+ {{{4}, {4}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
124+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
125+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
126+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
127+ },
128+ {
129+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
130+ },
131+ {
132+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
133+ });
134+ std::vector<std::vector<int64_t>> expectOutputShape = {{2, 8, 32, 32}};
135+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
136+}
137+ 
138+// case 4: 5D 输出,覆盖 tiling UT 中出现的最大维度场景
139+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_5d)
140+{
141+ vector<int64_t> shapeValue = {2, 3, 4, 5, 6};
142+ uint64_t keyValue = 42;
143+ uint64_t counterValue[2] = {0, 0};
144+ int32_t algValue = 1;
145+ gert::InfershapeContextPara infershapeContextPara(
146+ "StatelessTruncatedNormalV2",
147+ {
148+ {{{5}, {5}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
149+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
150+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
151+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
152+ },
153+ {
154+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
155+ },
156+ {
157+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
158+ });
159+ std::vector<std::vector<int64_t>> expectOutputShape = {{2, 3, 4, 5, 6}};
160+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
161+}
162+ 
163+// case 5: 含 1 的维度,确认不会被误当作标量或被压缩掉
164+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_dim_with_one)
165+{
166+ vector<int64_t> shapeValue = {1, 4097, 1};
167+ uint64_t keyValue = 42;
168+ uint64_t counterValue[2] = {0, 0};
169+ int32_t algValue = 1;
170+ gert::InfershapeContextPara infershapeContextPara(
171+ "StatelessTruncatedNormalV2",
172+ {
173+ {{{3}, {3}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
174+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
175+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
176+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
177+ },
178+ {
179+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
180+ },
181+ {
182+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
183+ });
184+ std::vector<std::vector<int64_t>> expectOutputShape = {{1, 4097, 1}};
185+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
186+}
187+ 
188+// case 6: float16 输出(dtype attr = 1,与本算子约定一致)
189+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_fp16)
190+{
191+ vector<int64_t> shapeValue = {64, 256};
192+ uint64_t keyValue = 42;
193+ uint64_t counterValue[2] = {0, 0};
194+ int32_t algValue = 1;
195+ gert::InfershapeContextPara infershapeContextPara(
196+ "StatelessTruncatedNormalV2",
197+ {
198+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
199+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
200+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
201+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
202+ },
203+ {
204+ {{{}, {}}, ge::DT_FLOAT16, ge::FORMAT_ND},
205+ },
206+ {
207+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(1)},
208+ });
209+ std::vector<std::vector<int64_t>> expectOutputShape = {{64, 256}};
210+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
211+}
212+ 
213+// case 7: bfloat16 输出(dtype attr = 27,与本算子约定一致)
214+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_bf16)
215+{
216+ vector<int64_t> shapeValue = {32, 512};
217+ uint64_t keyValue = 42;
218+ uint64_t counterValue[2] = {0, 0};
219+ int32_t algValue = 1;
220+ gert::InfershapeContextPara infershapeContextPara(
221+ "StatelessTruncatedNormalV2",
222+ {
223+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
224+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
225+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
226+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
227+ },
228+ {
229+ {{{}, {}}, ge::DT_BF16, ge::FORMAT_ND},
230+ },
231+ {
232+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(27)},
233+ });
234+ std::vector<std::vector<int64_t>> expectOutputShape = {{32, 512}};
235+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
236+}
237+ 
238+// case 8: INT32 shape tensor → DependencyMode 的 HandleShapeTensor<int32_t> 分支
239+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_int32_shape)
240+{
241+ vector<int32_t> shapeValue = {16, 64};
242+ uint64_t keyValue = 42;
243+ uint64_t counterValue[2] = {0, 0};
244+ int32_t algValue = 1;
245+ gert::InfershapeContextPara infershapeContextPara(
246+ "StatelessTruncatedNormalV2",
247+ {
248+ {{{2}, {2}}, ge::DT_INT32, ge::FORMAT_ND, true, shapeValue.data()},
249+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
250+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
251+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
252+ },
253+ {
254+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
255+ },
256+ {
257+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
258+ });
259+ std::vector<std::vector<int64_t>> expectOutputShape = {{16, 64}};
260+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
261+}
262+ 
263+// case 9: INT64 shape tensor(显式覆盖,与 case_1d 等共用分支但语义上独立列出)
264+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_int64_shape)
265+{
266+ vector<int64_t> shapeValue = {8, 128};
267+ uint64_t keyValue = 42;
268+ uint64_t counterValue[2] = {0, 0};
269+ int32_t algValue = 1;
270+ gert::InfershapeContextPara infershapeContextPara(
271+ "StatelessTruncatedNormalV2",
272+ {
273+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
274+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
275+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
276+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
277+ },
278+ {
279+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
280+ },
281+ {
282+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
283+ });
284+ std::vector<std::vector<int64_t>> expectOutputShape = {{8, 128}};
285+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
286+}
287+ 
288+// case 10: const data 缺失(1 元素)→ GetData 返回 nullptr → SetUnknownShape(1) → {-1}
289+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_null_const_1d)
290+{
291+ vector<int64_t> shapeValue = {};
292+ uint64_t keyValue = 42;
293+ uint64_t counterValue[2] = {0, 0};
294+ int32_t algValue = 1;
295+ gert::InfershapeContextPara infershapeContextPara(
296+ "StatelessTruncatedNormalV2",
297+ {
298+ {{{1}, {1}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
299+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
300+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
301+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
302+ },
303+ {
304+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
305+ },
306+ {
307+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
308+ });
309+ std::vector<std::vector<int64_t>> expectOutputShape = {{-1}};
310+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
311+}
312+ 
313+// case 11: const data 缺失,shape tensor 有 2 个元素 → SetUnknownShape(2) → {-1,-1}
314+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_null_const_2d)
315+{
316+ vector<int64_t> shapeValue = {};
317+ uint64_t keyValue = 42;
318+ uint64_t counterValue[2] = {0, 0};
319+ int32_t algValue = 1;
320+ gert::InfershapeContextPara infershapeContextPara(
321+ "StatelessTruncatedNormalV2",
322+ {
323+ {{{2}, {2}}, ge::DT_INT64, ge::FORMAT_ND, true, shapeValue.data()},
324+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
325+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
326+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
327+ },
328+ {
329+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
330+ },
331+ {
332+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
333+ });
334+ std::vector<std::vector<int64_t>> expectOutputShape = {{-1, -1}};
335+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
336+}
337+ 
338+// case 12: shape tensor 0 元素(标量输出场景)→ 输出 rank 0
339+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_0dim_scalar)
340+{
341+ int32_t shapeDummy = 0;
342+ uint64_t keyValue = 42;
343+ uint64_t counterValue[2] = {0, 0};
344+ int32_t algValue = 1;
345+ gert::InfershapeContextPara infershapeContextPara(
346+ "StatelessTruncatedNormalV2",
347+ {
348+ {{{0}, {0}}, ge::DT_INT32, ge::FORMAT_ND, true, &shapeDummy},
349+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
350+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
351+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
352+ },
353+ {
354+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
355+ },
356+ {
357+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
358+ });
359+ std::vector<std::vector<int64_t>> expectOutputShape = {{}};
360+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_SUCCESS, expectOutputShape);
361+}
362+ 
363+// case 13: 非法 shape dtype(DT_FLOAT)→ DependencyMode 两个分支都不命中 → GRAPH_FAILED
364+TEST_F(StatelessTruncatedNormalV2InferShapeTest, stateless_truncated_normal_v2_infershape_case_invalid_shape_dtype)
365+{
366+ vector<float> shapeValue = {32.0f, 512.0f};
367+ uint64_t keyValue = 42;
368+ uint64_t counterValue[2] = {0, 0};
369+ int32_t algValue = 1;
370+ gert::InfershapeContextPara infershapeContextPara(
371+ "StatelessTruncatedNormalV2",
372+ {
373+ {{{2}, {2}}, ge::DT_FLOAT, ge::FORMAT_ND, true, shapeValue.data()},
374+ {{{1}, {1}}, ge::DT_UINT64, ge::FORMAT_ND, true, &keyValue},
375+ {{{2}, {2}}, ge::DT_UINT64, ge::FORMAT_ND, true, counterValue},
376+ {{{}, {}}, ge::DT_INT32, ge::FORMAT_ND, true, &algValue},
377+ },
378+ {
379+ {{{}, {}}, ge::DT_FLOAT, ge::FORMAT_ND},
380+ },
381+ {
382+ {"dtype", Ops::Math::AnyValue::CreateFrom<int64_t>(0)},
383+ });
384+ ExecuteTestCase(infershapeContextPara, ge::GRAPH_FAILED);
385+}