已合并
fix: prevent unsupported data types from migrating to HostCPU #4523
ZhuXincheng创建于 18 天前
fix: prevent unsupported data types from migrating to HostCPU #4523
已合并
ZhuXincheng创建于 18 天前
3 个文件变更+70-6
@@ -45,6 +45,32 @@ const std::set<std::string> kControlV2Types = {IF, STATELESSIF, CASE, STATELESSC
45const std::set<std::string> kBlackList = {"RandomUniform"};45const std::set<std::string> kBlackList = {"RandomUniform"};
46const std::set<std::string> kWhiteList = {"MapIndex"};46const 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+ 
48bool IsGelocalOp(const OpDescPtr &op_desc) {74bool 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+ 
497TEST_F(UtestHostcpuEngineUpdatePass, HostInputWithoutConsumerNotMarkedAsModelInput) {528TEST_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();