已合并
feat: 增加RmsNormDynamicMxQuant infershape UT覆盖率 #5495
raoliang_sac创建于 5月30日
feat: 增加RmsNormDynamicMxQuant infershape UT覆盖率 #5495
已合并
raoliang_sac创建于 5月30日
共 1 个文件变更+353-0
@@ -229,4 +229,357 @@ TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_m_0_n_n
229TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_m_0_n_0_true)229TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_m_0_n_0_true)
230{230{
231 CheckInferShapeWithRstd({0, 0}, {0}, {0, 0}, {0, 0, 2}, {0, 1});231 CheckInferShapeWithRstd({0, 0}, {0}, {0, 0}, {0, 0, 2}, {0, 1});
232+}
233+ 
234+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_unknown_rank)
235+{
236+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
237+ auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_shape;
238+ 
239+ if (infer_shape_func != nullptr) {
240+ gert::StorageShape x_shape = {{-2}, {-2}};
241+ gert::StorageShape y_shape = {{-2}, {-2}};
242+ gert::StorageShape mxscale_shape = {{-2}, {-2}};
243+ gert::StorageShape rstd_shape = {{-2}, {-2}};
244+ 
245+ auto holder = gert::InferShapeContextFaker()
246+ .NodeIoNum(3, 3)
247+ .IrInstanceNum({1, 1, 1})
248+ .InputShapes({&x_shape, &x_shape, &x_shape})
249+ .OutputShapes({&y_shape, &mxscale_shape, &rstd_shape})
250+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)},
251+ {"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)},
252+ {"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")},
253+ {"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(40)},
254+ {"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
255+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
256+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
257+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
258+ .NodeOutputTd(0, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND, ge::FORMAT_ND)
259+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
260+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
261+ .Build();
262+ 
263+ auto context = holder.GetContext<gert::InferShapeContext>();
264+ EXPECT_EQ(infer_shape_func(context), ge::GRAPH_SUCCESS);
265+ 
266+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(0), -2);
267+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(0), -2);
268+ }
269+}
270+ 
271+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_unknown_rank_with_rstd)
272+{
273+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
274+ auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_shape;
275+ 
276+ if (infer_shape_func != nullptr) {
277+ gert::StorageShape x_shape = {{-2}, {-2}};
278+ gert::StorageShape y_shape = {{-2}, {-2}};
279+ gert::StorageShape mxscale_shape = {{-2}, {-2}};
280+ gert::StorageShape rstd_shape = {{-2}, {-2}};
281+ 
282+ auto holder = gert::InferShapeContextFaker()
283+ .NodeIoNum(3, 3)
284+ .IrInstanceNum({1, 1, 1})
285+ .InputShapes({&x_shape, &x_shape, &x_shape})
286+ .OutputShapes({&y_shape, &mxscale_shape, &rstd_shape})
287+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)},
288+ {"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)},
289+ {"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")},
290+ {"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(40)},
291+ {"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(true)}})
292+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
293+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
294+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
295+ .NodeOutputTd(0, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND, ge::FORMAT_ND)
296+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
297+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
298+ .Build();
299+ 
300+ auto context = holder.GetContext<gert::InferShapeContext>();
301+ EXPECT_EQ(infer_shape_func(context), ge::GRAPH_SUCCESS);
302+ 
303+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(0), -2);
304+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(0), -2);
305+ EXPECT_EQ(context->GetOutputShape(2)->GetDim(0), -2);
306+ }
307+}
308+ 
309+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_unknown_dim)
310+{
311+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
312+ auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_shape;
313+ 
314+ if (infer_shape_func != nullptr) {
315+ gert::StorageShape x_shape = {{8, -1}, {8, -1}};
316+ gert::StorageShape gamma_shape = {{-1}, {-1}};
317+ gert::StorageShape y_shape = {{8, -1}, {8, -1}};
318+ gert::StorageShape mxscale_shape = {{8, -1, 2}, {8, -1, 2}};
319+ gert::StorageShape rstd_shape = {{8, 1}, {8, 1}};
320+ 
321+ auto holder = gert::InferShapeContextFaker()
322+ .NodeIoNum(3, 3)
323+ .IrInstanceNum({1, 1, 1})
324+ .InputShapes({&x_shape, &gamma_shape, &gamma_shape})
325+ .OutputShapes({&y_shape, &mxscale_shape, &rstd_shape})
326+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)},
327+ {"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)},
328+ {"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")},
329+ {"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(40)},
330+ {"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
331+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
332+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
333+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
334+ .NodeOutputTd(0, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND, ge::FORMAT_ND)
335+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
336+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
337+ .Build();
338+ 
339+ auto context = holder.GetContext<gert::InferShapeContext>();
340+ EXPECT_EQ(infer_shape_func(context), ge::GRAPH_SUCCESS);
341+ 
342+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(0), 8);
343+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(1), -1);
344+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(0), 8);
345+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(1), -1);
346+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(2), 2);
347+ }
348+}
349+ 
350+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_unknown_dim_with_rstd)
351+{
352+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
353+ auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_shape;
354+ 
355+ if (infer_shape_func != nullptr) {
356+ gert::StorageShape x_shape = {{8, -1}, {8, -1}};
357+ gert::StorageShape gamma_shape = {{-1}, {-1}};
358+ gert::StorageShape y_shape = {{8, -1}, {8, -1}};
359+ gert::StorageShape mxscale_shape = {{8, -1, 2}, {8, -1, 2}};
360+ gert::StorageShape rstd_shape = {{8, 1}, {8, 1}};
361+ 
362+ auto holder = gert::InferShapeContextFaker()
363+ .NodeIoNum(3, 3)
364+ .IrInstanceNum({1, 1, 1})
365+ .InputShapes({&x_shape, &gamma_shape, &gamma_shape})
366+ .OutputShapes({&y_shape, &mxscale_shape, &rstd_shape})
367+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)},
368+ {"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)},
369+ {"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")},
370+ {"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(40)},
371+ {"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(true)}})
372+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
373+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
374+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
375+ .NodeOutputTd(0, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND, ge::FORMAT_ND)
376+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
377+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
378+ .Build();
379+ 
380+ auto context = holder.GetContext<gert::InferShapeContext>();
381+ EXPECT_EQ(infer_shape_func(context), ge::GRAPH_SUCCESS);
382+ 
383+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(0), 8);
384+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(1), -1);
385+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(0), 8);
386+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(1), -1);
387+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(2), 2);
388+ EXPECT_EQ(context->GetOutputShape(2)->GetDim(0), 8);
389+ EXPECT_EQ(context->GetOutputShape(2)->GetDim(1), 1);
390+ }
391+}
392+ 
393+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infer_dtype_float4_e2m1)
394+{
395+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
396+ auto data_type_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_datatype;
397+ 
398+ if (data_type_func != nullptr) {
399+ ge::DataType input_ref = ge::DT_FLOAT16;
400+ ge::DataType y_ref = ge::DT_FLOAT4_E2M1;
401+ ge::DataType mx_scale_ref = ge::DT_FLOAT8_E8M0;
402+ ge::DataType rstd_ref = ge::DT_FLOAT;
403+ auto context_holder =
404+ gert::InferDataTypeContextFaker()
405+ .IrInputNum(3)
406+ .NodeIoNum(3, 3)
407+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
408+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
409+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
410+ .NodeOutputTd(0, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND, ge::FORMAT_ND)
411+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
412+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
413+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)}})
414+ .NodeAttrs({{"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)}})
415+ .NodeAttrs({{"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")}})
416+ .NodeAttrs({{"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(40)}})
417+ .NodeAttrs({{"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
418+ .InputDataTypes({&input_ref, &input_ref, &input_ref})
419+ .OutputDataTypes({&y_ref, &mx_scale_ref, &rstd_ref})
420+ .Build();
421+ auto context = context_holder.GetContext<gert::InferDataTypeContext>();
422+ EXPECT_EQ(data_type_func(context), ge::GRAPH_SUCCESS);
423+ ASSERT_NE(context, nullptr);
424+ 
425+ EXPECT_EQ(context->GetOutputDataType(0), y_ref);
426+ EXPECT_EQ(context->GetOutputDataType(1), mx_scale_ref);
427+ EXPECT_EQ(context->GetOutputDataType(2), rstd_ref);
428+ }
429+}
430+ 
431+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infer_dtype_float4_e1m2)
432+{
433+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
434+ auto data_type_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_datatype;
435+ 
436+ if (data_type_func != nullptr) {
437+ ge::DataType input_ref = ge::DT_FLOAT16;
438+ ge::DataType y_ref = ge::DT_FLOAT4_E1M2;
439+ ge::DataType mx_scale_ref = ge::DT_FLOAT8_E8M0;
440+ ge::DataType rstd_ref = ge::DT_FLOAT;
441+ auto context_holder =
442+ gert::InferDataTypeContextFaker()
443+ .IrInputNum(3)
444+ .NodeIoNum(3, 3)
445+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
446+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
447+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
448+ .NodeOutputTd(0, ge::DT_FLOAT4_E1M2, ge::FORMAT_ND, ge::FORMAT_ND)
449+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
450+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
451+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)}})
452+ .NodeAttrs({{"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)}})
453+ .NodeAttrs({{"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")}})
454+ .NodeAttrs({{"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(41)}})
455+ .NodeAttrs({{"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
456+ .InputDataTypes({&input_ref, &input_ref, &input_ref})
457+ .OutputDataTypes({&y_ref, &mx_scale_ref, &rstd_ref})
458+ .Build();
459+ auto context = context_holder.GetContext<gert::InferDataTypeContext>();
460+ EXPECT_EQ(data_type_func(context), ge::GRAPH_SUCCESS);
461+ ASSERT_NE(context, nullptr);
462+ 
463+ EXPECT_EQ(context->GetOutputDataType(0), y_ref);
464+ EXPECT_EQ(context->GetOutputDataType(1), mx_scale_ref);
465+ EXPECT_EQ(context->GetOutputDataType(2), rstd_ref);
466+ }
467+}
468+ 
469+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infer_dtype_float8_e4m3fn)
470+{
471+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
472+ auto data_type_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_datatype;
473+ 
474+ if (data_type_func != nullptr) {
475+ ge::DataType input_ref = ge::DT_FLOAT16;
476+ ge::DataType y_ref = ge::DT_FLOAT8_E4M3FN;
477+ ge::DataType mx_scale_ref = ge::DT_FLOAT8_E8M0;
478+ ge::DataType rstd_ref = ge::DT_FLOAT;
479+ auto context_holder =
480+ gert::InferDataTypeContextFaker()
481+ .IrInputNum(3)
482+ .NodeIoNum(3, 3)
483+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
484+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
485+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
486+ .NodeOutputTd(0, ge::DT_FLOAT8_E4M3FN, ge::FORMAT_ND, ge::FORMAT_ND)
487+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
488+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
489+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)}})
490+ .NodeAttrs({{"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)}})
491+ .NodeAttrs({{"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")}})
492+ .NodeAttrs({{"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(36)}})
493+ .NodeAttrs({{"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
494+ .InputDataTypes({&input_ref, &input_ref, &input_ref})
495+ .OutputDataTypes({&y_ref, &mx_scale_ref, &rstd_ref})
496+ .Build();
497+ auto context = context_holder.GetContext<gert::InferDataTypeContext>();
498+ EXPECT_EQ(data_type_func(context), ge::GRAPH_SUCCESS);
499+ ASSERT_NE(context, nullptr);
500+ 
501+ EXPECT_EQ(context->GetOutputDataType(0), y_ref);
502+ EXPECT_EQ(context->GetOutputDataType(1), mx_scale_ref);
503+ EXPECT_EQ(context->GetOutputDataType(2), rstd_ref);
504+ }
505+}
506+ 
507+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infer_dtype_invalid)
508+{
509+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
510+ auto data_type_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_datatype;
511+ 
512+ if (data_type_func != nullptr) {
513+ ge::DataType input_ref = ge::DT_FLOAT16;
514+ ge::DataType y_ref = ge::DT_FLOAT16;
515+ ge::DataType mx_scale_ref = ge::DT_FLOAT8_E8M0;
516+ ge::DataType rstd_ref = ge::DT_FLOAT;
517+ auto context_holder =
518+ gert::InferDataTypeContextFaker()
519+ .IrInputNum(3)
520+ .NodeIoNum(3, 3)
521+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
522+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
523+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
524+ .NodeOutputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
525+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
526+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
527+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)}})
528+ .NodeAttrs({{"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)}})
529+ .NodeAttrs({{"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")}})
530+ .NodeAttrs({{"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(999)}})
531+ .NodeAttrs({{"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
532+ .InputDataTypes({&input_ref, &input_ref, &input_ref})
533+ .OutputDataTypes({&y_ref, &mx_scale_ref, &rstd_ref})
534+ .Build();
535+ auto context = context_holder.GetContext<gert::InferDataTypeContext>();
536+ EXPECT_EQ(data_type_func(context), ge::GRAPH_FAILED);
537+ }
538+}
539+ 
540+TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_debug_log)
541+{
542+ setenv("ASCEND_GLOBAL_LOG_LEVEL", "0", true);
543+ dlog_setlevel(-1, 0, 1);
544+ 
545+ ASSERT_NE(gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant"), nullptr);
546+ auto infer_shape_func = gert::OpImplRegistry::GetInstance().GetOpImpl("RmsNormDynamicMxQuant")->infer_shape;
547+ 
548+ ASSERT_NE(infer_shape_func, nullptr);
549+ 
550+ gert::StorageShape x_shape = {{8, 64}, {8, 64}};
551+ gert::StorageShape gamma_shape = {{64}, {64}};
552+ gert::StorageShape y_shape = {{8, 64}, {8, 64}};
553+ gert::StorageShape mxscale_shape = {{8, 1, 2}, {8, 1, 2}};
554+ gert::StorageShape rstd_shape = {{8, 1}, {8, 1}};
555+ 
556+ auto holder = gert::InferShapeContextFaker()
557+ .NodeIoNum(3, 3)
558+ .IrInstanceNum({1, 1, 1})
559+ .InputShapes({&x_shape, &gamma_shape, &gamma_shape})
560+ .OutputShapes({&y_shape, &mxscale_shape, &rstd_shape})
561+ .NodeAttrs({{"epsilon", Ops::NN::AnyValue::CreateFrom<float>(1e-6)},
562+ {"scale_alg", Ops::NN::AnyValue::CreateFrom<int64_t>(0)},
563+ {"round_mode", Ops::NN::AnyValue::CreateFrom<string>("rint")},
564+ {"dst_type", Ops::NN::AnyValue::CreateFrom<int64_t>(40)},
565+ {"output_rstd", Ops::NN::AnyValue::CreateFrom<bool>(false)}})
566+ .NodeInputTd(0, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
567+ .NodeInputTd(1, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
568+ .NodeInputTd(2, ge::DT_FLOAT16, ge::FORMAT_ND, ge::FORMAT_ND)
569+ .NodeOutputTd(0, ge::DT_FLOAT4_E2M1, ge::FORMAT_ND, ge::FORMAT_ND)
570+ .NodeOutputTd(1, ge::DT_FLOAT8_E8M0, ge::FORMAT_ND, ge::FORMAT_ND)
571+ .NodeOutputTd(2, ge::DT_FLOAT, ge::FORMAT_ND, ge::FORMAT_ND)
572+ .Build();
573+ 
574+ auto context = holder.GetContext<gert::InferShapeContext>();
575+ EXPECT_EQ(infer_shape_func(context), ge::GRAPH_SUCCESS);
576+ 
577+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(0), 8);
578+ EXPECT_EQ(context->GetOutputShape(0)->GetDim(1), 64);
579+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(0), 8);
580+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(1), 1);
581+ EXPECT_EQ(context->GetOutputShape(1)->GetDim(2), 2);
582+ 
583+ dlog_setlevel(-1, 2, 1);
584+ setenv("ASCEND_GLOBAL_LOG_LEVEL", "3", true);
232}585}