已合并
feat: 增加RmsNormDynamicMxQuant infershape UT覆盖率 #5495
raoliang_sac创建于 5月30日
feat: 增加RmsNormDynamicMxQuant infershape UT覆盖率 #5495
已合并
共 1 个文件变更+353-0
Mnorm/rms_norm_dynamic_mx_quant/tests/ut/op_host/test_rms_norm_dynamic_mx_quant_infershape.cpp+353-0
| @@ -229,4 +229,357 @@ TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_m_0_n_n | |||
| 229 | TEST_F(RmsNormDynamicMxQuantInfershape, RmsNormDynamicMxQuant_infershape_m_0_n_0_true) | 229 | TEST_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 | } |