已合并
fix: prevent unsupported data types from migrating to HostCPU #4523
ZhuXincheng创建于 18 天前
fix: prevent unsupported data types from migrating to HostCPU #4523
已合并
共 3 个文件变更+70-6
| @@ -45,6 +45,32 @@ const std::set<std::string> kControlV2Types = {IF, STATELESSIF, CASE, STATELESSC | |||
| 45 | const std::set<std::string> kBlackList = {"RandomUniform"}; | 45 | const std::set<std::string> kBlackList = {"RandomUniform"}; |
| 46 | const std::set<std::string> kWhiteList = {"MapIndex"}; | 46 | const std::set<std::string> kWhiteList = {"MapIndex"}; |
| 47 | 47 | ||
| 48 | +bool IsUnsupportedHostCpuDataType(const DataType data_type) { | ||
| 49 | + switch (data_type) { | ||
| 50 | + case DT_STRING: | ||
| 51 | + case DT_STRING_REF: | ||
| 52 | + case DT_VARIANT: | ||
| 53 | + case DT_RESOURCE: | ||
| 54 | + return true; | ||
| 55 | + default: | ||
| 56 | + return false; | ||
| 57 | + } | ||
| 58 | +} | ||
| 59 | + | ||
| 60 | +bool HasUnsupportedHostCpuDataType(const OpDescPtr &op_desc) { | ||
| 61 | + for (const auto &tensor_desc : op_desc->GetAllInputsDesc()) { | ||
| 62 | + if (IsUnsupportedHostCpuDataType(tensor_desc.GetDataType())) { | ||
| 63 | + return true; | ||
| 64 | + } | ||
| 65 | + } | ||
| 66 | + for (const auto &tensor_desc : op_desc->GetAllOutputsDesc()) { | ||
| 67 | + if (IsUnsupportedHostCpuDataType(tensor_desc.GetDataType())) { | ||
| 68 | + return true; | ||
| 69 | + } | ||
| 70 | + } | ||
| 71 | + return false; | ||
| 72 | +} | ||
| 73 | + | ||
| 48 | bool IsGelocalOp(const OpDescPtr &op_desc) { | 74 | bool IsGelocalOp(const OpDescPtr &op_desc) { |
| 49 | return op_desc->GetOpKernelLibName() == kGeLocalOpKernelLibName; | 75 | return op_desc->GetOpKernelLibName() == kGeLocalOpKernelLibName; |
| 50 | } | 76 | } |
| @@ -126,6 +152,11 @@ bool IsSupportHostcpu(const OpDescPtr &op_desc) { | |||
| 126 | if (kBlackList.count(op_desc->GetType()) > 0U) { | 152 | if (kBlackList.count(op_desc->GetType()) > 0U) { |
| 127 | return false; | 153 | return false; |
| 128 | } | 154 | } |
| 155 | + if (HasUnsupportedHostCpuDataType(op_desc)) { | ||
| 156 | + GELOGI("[HostcpuEngineUpdatePass]: host cpu does not support node[%s] type[%s] due to unsupported data type.", | ||
| 157 | + op_desc->GetName().c_str(), op_desc->GetType().c_str()); | ||
| 158 | + return false; | ||
| 159 | + } | ||
| 129 | auto op_infos = OpsKernelManager::GetInstance().GetOpsKernelInfo(op_desc->GetType()); | 160 | auto op_infos = OpsKernelManager::GetInstance().GetOpsKernelInfo(op_desc->GetType()); |
| 130 | for (const auto &it : op_infos) { | 161 | for (const auto &it : op_infos) { |
| 131 | if ((it.engine == kHostCpuEngineName) && (it.opKernelLib == kHostCpuOpKernelLibName)) { | 162 | if ((it.engine == kHostCpuEngineName) && (it.opKernelLib == kHostCpuOpKernelLibName)) { |
| @@ -1997,7 +1997,9 @@ TEST_F(DynamicGraphTest, HostCpuPassDoesNotRouteUnsupportedConcatV2ToHostCpu) { | |||
| 1997 | NodeEngineMap node_composite_engine_map; | 1997 | NodeEngineMap node_composite_engine_map; |
| 1998 | EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS); | 1998 | EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS); |
| 1999 | 1999 | ||
| 2000 | - EXPECT_GT(host_cpu_pass_env.GetCheckSupportedCount(), 0U); | 2000 | + // Unsupported data types are rejected by HostcpuEngineUpdatePass before |
| 2001 | + // consulting the HostCPU kernel store. | ||
| 2002 | + EXPECT_EQ(host_cpu_pass_env.GetCheckSupportedCount(), 0U); | ||
| 2001 | EXPECT_EQ(concat->GetOpDesc()->GetOpKernelLibName(), kEngineNameAiCore); | 2003 | EXPECT_EQ(concat->GetOpDesc()->GetOpKernelLibName(), kEngineNameAiCore); |
| 2002 | EXPECT_TRUE(node_atomic_engine_map.count(concat) == 0U); | 2004 | EXPECT_TRUE(node_atomic_engine_map.count(concat) == 0U); |
| 2003 | bool is_host_model_input = false; | 2005 | bool is_host_model_input = false; |
| @@ -2006,7 +2008,7 @@ TEST_F(DynamicGraphTest, HostCpuPassDoesNotRouteUnsupportedConcatV2ToHostCpu) { | |||
| 2006 | unsetenv("ENABLE_RUNTIME_V2"); | 2008 | unsetenv("ENABLE_RUNTIME_V2"); |
| 2007 | } | 2009 | } |
| 2008 | 2010 | ||
| 2009 | -TEST_F(DynamicGraphTest, HostCpuPassRoutesConcatV2WhenHostCpuStoreMissing) { | 2011 | +TEST_F(DynamicGraphTest, HostCpuPassDoesNotRouteConcatV2WhenHostCpuStoreMissing) { |
| 2010 | setenv("ENABLE_RUNTIME_V2", "1", 1); | 2012 | setenv("ENABLE_RUNTIME_V2", "1", 1); |
| 2011 | ScopedHostCpuPassEnv host_cpu_pass_env(false); | 2013 | ScopedHostCpuPassEnv host_cpu_pass_env(false); |
| 2012 | auto graph = BuildHostCpuUnsupportedConcatV2Graph(); | 2014 | auto graph = BuildHostCpuUnsupportedConcatV2Graph(); |
| @@ -2021,12 +2023,12 @@ TEST_F(DynamicGraphTest, HostCpuPassRoutesConcatV2WhenHostCpuStoreMissing) { | |||
| 2021 | NodeEngineMap node_composite_engine_map; | 2023 | NodeEngineMap node_composite_engine_map; |
| 2022 | EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS); | 2024 | EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS); |
| 2023 | 2025 | ||
| 2024 | - EXPECT_EQ(concat->GetOpDesc()->GetOpKernelLibName(), kStHostCpuKernelStore); | 2026 | + EXPECT_EQ(concat->GetOpDesc()->GetOpKernelLibName(), kEngineNameAiCore); |
| 2025 | - EXPECT_EQ(concat->GetOpDesc()->GetOpEngineName(), kStHostCpuEngine); | 2027 | + EXPECT_TRUE(concat->GetOpDesc()->GetOpEngineName().empty()); |
| 2026 | - EXPECT_EQ(node_atomic_engine_map[concat], kStHostCpuEngine); | 2028 | + EXPECT_TRUE(node_atomic_engine_map.count(concat) == 0U); |
| 2027 | bool is_host_model_input = false; | 2029 | bool is_host_model_input = false; |
| 2028 | (void)AttrUtils::GetBool(data->GetOpDesc(), ATTR_NAME_HOST_TENSOR_AS_MODEL_INPUT, is_host_model_input); | 2030 | (void)AttrUtils::GetBool(data->GetOpDesc(), ATTR_NAME_HOST_TENSOR_AS_MODEL_INPUT, is_host_model_input); |
| 2029 | - EXPECT_TRUE(is_host_model_input); | 2031 | + EXPECT_FALSE(is_host_model_input); |
| 2030 | unsetenv("ENABLE_RUNTIME_V2"); | 2032 | unsetenv("ENABLE_RUNTIME_V2"); |
| 2031 | } | 2033 | } |
| 2032 | 2034 | ||
| @@ -494,6 +494,37 @@ TEST_F(UtestHostcpuEngineUpdatePass, HostCpuSupportCheckFailsWithKernelStore) { | |||
| 494 | EXPECT_TRUE(node_atomic_engine_map.count(graph->FindNode("gather")) == 0U); | 494 | EXPECT_TRUE(node_atomic_engine_map.count(graph->FindNode("gather")) == 0U); |
| 495 | } | 495 | } |
| 496 | 496 | ||
| 497 | +TEST_F(UtestHostcpuEngineUpdatePass, HostCpuRejectsUnsupportedDataTypes) { | ||
| 498 | + setenv("ENABLE_RUNTIME_V2", "1", 1); | ||
| 499 | + const std::vector<DataType> unsupported_types = {DT_STRING, DT_STRING_REF, DT_VARIANT, DT_RESOURCE}; | ||
| 500 | + for (const auto data_type : unsupported_types) { | ||
| 501 | + auto graph = BuildHostCpuSupportCheckGraph(); | ||
| 502 | + ASSERT_NE(graph, nullptr); | ||
| 503 | + auto gather_desc = graph->FindNode("gather")->GetOpDesc(); | ||
| 504 | + gather_desc->MutableInputDesc(0)->SetDataType(data_type); | ||
| 505 | + gather_desc->MutableInputDesc(0)->SetOriginDataType(data_type); | ||
| 506 | + | ||
| 507 | + HostcpuEngineUpdatePass pass; | ||
| 508 | + NodeEngineMap node_atomic_engine_map; | ||
| 509 | + NodeEngineMap node_composite_engine_map; | ||
| 510 | + EXPECT_EQ(pass.Run(graph, node_atomic_engine_map, node_composite_engine_map), SUCCESS); | ||
| 511 | + EXPECT_EQ(gather_desc->GetOpKernelLibName(), kEngineNameAiCore); | ||
| 512 | + EXPECT_TRUE(node_atomic_engine_map.count(graph->FindNode("gather")) == 0U); | ||
| 513 | + } | ||
| 514 | + | ||
| 515 | + auto output_graph = BuildHostCpuSupportCheckGraph(); | ||
| 516 | + ASSERT_NE(output_graph, nullptr); | ||
| 517 | + auto output_gather_desc = output_graph->FindNode("gather")->GetOpDesc(); | ||
| 518 | + output_gather_desc->MutableOutputDesc(0)->SetDataType(DT_STRING); | ||
| 519 | + output_gather_desc->MutableOutputDesc(0)->SetOriginDataType(DT_STRING); | ||
| 520 | + HostcpuEngineUpdatePass output_pass; | ||
| 521 | + NodeEngineMap output_atomic_engine_map; | ||
| 522 | + NodeEngineMap output_composite_engine_map; | ||
| 523 | + EXPECT_EQ(output_pass.Run(output_graph, output_atomic_engine_map, output_composite_engine_map), SUCCESS); | ||
| 524 | + EXPECT_EQ(output_gather_desc->GetOpKernelLibName(), kEngineNameAiCore); | ||
| 525 | + EXPECT_TRUE(output_atomic_engine_map.count(output_graph->FindNode("gather")) == 0U); | ||
| 526 | +} | ||
| 527 | + | ||
| 497 | TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInput) { | 528 | TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInput) { |
| 498 | setenv("ENABLE_RUNTIME_V2", "1", 1); | 529 | setenv("ENABLE_RUNTIME_V2", "1", 1); |
| 499 | auto graph = BuildHostInputWithoutConsumerGraph(); | 530 | auto graph = BuildHostInputWithoutConsumerGraph(); |